lee-t/ControlNet-Video
0
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 huggingface_hub import hf_hub_url16from pytorch_lightning import seed_everything17 18sys.path.append('ControlNet')19 20import config21from annotator.canny import apply_canny22from annotator.hed import apply_hed, nms23from annotator.midas import apply_midas24from annotator.mlsd import apply_mlsd25from annotator.openpose import apply_openpose26from annotator.uniformer import apply_uniformer27from annotator.util import HWC3, resize_image28from cldm.model import create_model, load_state_dict29from ldm.models.diffusion.ddim import DDIMSampler30from share import *31 32 33MODEL_NAMES = {34 'canny': 'control_canny-fp16.safetensors',35 'hough': 'control_mlsd-fp16.safetensors',36 'hed': 'control_hed-fp16.safetensors',37 'scribble': 'control_scribble-fp16.safetensors',38 'pose': 'control_openpose-fp16.safetensors',39 'seg': 'control_seg-fp16.safetensors',40 'depth': 'control_depth-fp16.safetensors',41 'normal': 'control_normal-fp16.safetensors',42}43 44MODEL_REPO = 'webui/ControlNet-modules-safetensors'45 46DEFAULT_BASE_MODEL_REPO = 'runwayml/stable-diffusion-v1-5'47DEFAULT_BASE_MODEL_FILENAME = 'v1-5-pruned-emaonly.safetensors'48DEFAULT_BASE_MODEL_URL = 'https://huggingface.co/runwayml/stable-diffusion-v1-5/resolve/main/v1-5-pruned-emaonly.safetensors'49 50class Model:51 def __init__(self,52 model_config_path: str = 'ControlNet/models/cldm_v15.yaml',53 model_dir: str = 'models'):54 self.device = torch.device(55 'cuda:0' if torch.cuda.is_available() else 'cpu')56 self.model = create_model(model_config_path).to(self.device)57 self.ddim_sampler = DDIMSampler(self.model)58 self.task_name = ''59 60 self.base_model_url = ''61 62 self.model_dir = pathlib.Path(model_dir)63 self.model_dir.mkdir(exist_ok=True, parents=True)64 65 self.download_models()66 self.set_base_model(DEFAULT_BASE_MODEL_REPO,67 DEFAULT_BASE_MODEL_FILENAME)68 69 def set_base_model(self, model_id: str, filename: str) -> str:70 if not model_id or not filename:71 return self.base_model_url72 base_model_url = hf_hub_url(model_id, filename)73 if base_model_url != self.base_model_url:74 self.load_base_model(base_model_url)75 self.base_model_url = base_model_url76 return self.base_model_url77 78 79 def download_base_model(self, model_url: str) -> pathlib.Path:80 self.model_dir.mkdir(exist_ok=True, parents=True)81 model_name = model_url.split('/')[-1]82 out_path = self.model_dir / model_name83 if not out_path.exists():84 subprocess.run(shlex.split(f'wget {model_url} -O {out_path}'))85 return out_path86 87 def load_base_model(self, model_url: str) -> None:88 model_path = self.download_base_model(model_url)89 self.model.load_state_dict(load_state_dict(model_path,90 location=self.device.type),91 strict=False)92 93 def load_weight(self, task_name: str) -> None:94 if task_name == self.task_name:95 return96 weight_path = self.get_weight_path(task_name)97 self.model.control_model.load_state_dict(98 load_state_dict(weight_path, location=self.device.type))99 self.task_name = task_name100 101 def get_weight_path(self, task_name: str) -> str:102 if 'scribble' in task_name:103 task_name = 'scribble'104 return f'{self.model_dir}/{MODEL_NAMES[task_name]}'105 106 def download_models(self) -> None:107 self.model_dir.mkdir(exist_ok=True, parents=True)108 for name in MODEL_NAMES.values():109 out_path = self.model_dir / name110 if out_path.exists():111 continue112 model_url = hf_hub_url(MODEL_REPO, name)113 subprocess.run(shlex.split(f'wget {model_url} -O {out_path}'))114 115 @torch.inference_mode()116 def process_canny(self, input_image, prompt, a_prompt, n_prompt,117 num_samples, image_resolution, ddim_steps, scale, seed,118 eta, low_threshold, high_threshold):119 self.load_weight('canny')120 121 img = resize_image(HWC3(input_image), image_resolution)122 H, W, C = img.shape123 124 detected_map = apply_canny(img, low_threshold, high_threshold)125 detected_map = HWC3(detected_map)126 127 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0128 control = torch.stack([control for _ in range(num_samples)], dim=0)129 control = einops.rearrange(control, 'b h w c -> b c h w').clone()130 131 if seed == -1:132 seed = random.randint(0, 65535)133 seed_everything(seed)134 135 if config.save_memory:136 self.model.low_vram_shift(is_diffusing=False)137 138 cond = {139 'c_concat': [control],140 'c_crossattn': [141 self.model.get_learned_conditioning(142 [prompt + ', ' + a_prompt] * num_samples)143 ]144 }145 un_cond = {146 'c_concat': [control],147 'c_crossattn':148 [self.model.get_learned_conditioning([n_prompt] * num_samples)]149 }150 shape = (4, H // 8, W // 8)151 152 if config.save_memory:153 self.model.low_vram_shift(is_diffusing=True)154 155 samples, intermediates = self.ddim_sampler.sample(156 ddim_steps,157 num_samples,158 shape,159 cond,160 verbose=False,161 eta=eta,162 unconditional_guidance_scale=scale,163 unconditional_conditioning=un_cond)164 165 if config.save_memory:166 self.model.low_vram_shift(is_diffusing=False)167 168 x_samples = self.model.decode_first_stage(samples)169 x_samples = (170 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +171 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)172 173 results = [x_samples[i] for i in range(num_samples)]174 return [255 - detected_map] + results175 176 @torch.inference_mode()177 def process_hough(self, input_image, prompt, a_prompt, n_prompt,178 num_samples, image_resolution, detect_resolution,179 ddim_steps, scale, seed, eta, value_threshold,180 distance_threshold):181 self.load_weight('hough')182 183 input_image = HWC3(input_image)184 detected_map = apply_mlsd(resize_image(input_image, detect_resolution),185 value_threshold, distance_threshold)186 detected_map = HWC3(detected_map)187 img = resize_image(input_image, image_resolution)188 H, W, C = img.shape189 190 detected_map = cv2.resize(detected_map, (W, H),191 interpolation=cv2.INTER_NEAREST)192 193 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0194 control = torch.stack([control for _ in range(num_samples)], dim=0)195 control = einops.rearrange(control, 'b h w c -> b c h w').clone()196 197 if seed == -1:198 seed = random.randint(0, 65535)199 seed_everything(seed)200 201 if config.save_memory:202 self.model.low_vram_shift(is_diffusing=False)203 204 cond = {205 'c_concat': [control],206 'c_crossattn': [207 self.model.get_learned_conditioning(208 [prompt + ', ' + a_prompt] * num_samples)209 ]210 }211 un_cond = {212 'c_concat': [control],213 'c_crossattn':214 [self.model.get_learned_conditioning([n_prompt] * num_samples)]215 }216 shape = (4, H // 8, W // 8)217 218 if config.save_memory:219 self.model.low_vram_shift(is_diffusing=True)220 221 samples, intermediates = self.ddim_sampler.sample(222 ddim_steps,223 num_samples,224 shape,225 cond,226 verbose=False,227 eta=eta,228 unconditional_guidance_scale=scale,229 unconditional_conditioning=un_cond)230 231 if config.save_memory:232 self.model.low_vram_shift(is_diffusing=False)233 234 x_samples = self.model.decode_first_stage(samples)235 x_samples = (236 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +237 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)238 239 results = [x_samples[i] for i in range(num_samples)]240 return [241 255 - cv2.dilate(detected_map,242 np.ones(shape=(3, 3), dtype=np.uint8),243 iterations=1)244 ] + results245 246 @torch.inference_mode()247 def process_hed(self, input_image, prompt, a_prompt, n_prompt, num_samples,248 image_resolution, detect_resolution, ddim_steps, scale,249 seed, eta):250 self.load_weight('hed')251 252 input_image = HWC3(input_image)253 detected_map = apply_hed(resize_image(input_image, detect_resolution))254 detected_map = HWC3(detected_map)255 img = resize_image(input_image, image_resolution)256 H, W, C = img.shape257 258 detected_map = cv2.resize(detected_map, (W, H),259 interpolation=cv2.INTER_LINEAR)260 261 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0262 control = torch.stack([control for _ in range(num_samples)], dim=0)263 control = einops.rearrange(control, 'b h w c -> b c h w').clone()264 265 if seed == -1:266 seed = random.randint(0, 65535)267 seed_everything(seed)268 269 if config.save_memory:270 self.model.low_vram_shift(is_diffusing=False)271 272 cond = {273 'c_concat': [control],274 'c_crossattn': [275 self.model.get_learned_conditioning(276 [prompt + ', ' + a_prompt] * num_samples)277 ]278 }279 un_cond = {280 'c_concat': [control],281 'c_crossattn':282 [self.model.get_learned_conditioning([n_prompt] * num_samples)]283 }284 shape = (4, H // 8, W // 8)285 286 if config.save_memory:287 self.model.low_vram_shift(is_diffusing=True)288 289 samples, intermediates = self.ddim_sampler.sample(290 ddim_steps,291 num_samples,292 shape,293 cond,294 verbose=False,295 eta=eta,296 unconditional_guidance_scale=scale,297 unconditional_conditioning=un_cond)298 299 if config.save_memory:300 self.model.low_vram_shift(is_diffusing=False)301 302 x_samples = self.model.decode_first_stage(samples)303 x_samples = (304 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +305 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)306 307 results = [x_samples[i] for i in range(num_samples)]308 return [detected_map] + results309 310 @torch.inference_mode()311 def process_scribble(self, input_image, prompt, a_prompt, n_prompt,312 num_samples, image_resolution, ddim_steps, scale,313 seed, eta):314 self.load_weight('scribble')315 316 img = resize_image(HWC3(input_image), image_resolution)317 H, W, C = img.shape318 319 detected_map = np.zeros_like(img, dtype=np.uint8)320 detected_map[np.min(img, axis=2) < 127] = 255321 322 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0323 control = torch.stack([control for _ in range(num_samples)], dim=0)324 control = einops.rearrange(control, 'b h w c -> b c h w').clone()325 326 if seed == -1:327 seed = random.randint(0, 65535)328 seed_everything(seed)329 330 if config.save_memory:331 self.model.low_vram_shift(is_diffusing=False)332 333 cond = {334 'c_concat': [control],335 'c_crossattn': [336 self.model.get_learned_conditioning(337 [prompt + ', ' + a_prompt] * num_samples)338 ]339 }340 un_cond = {341 'c_concat': [control],342 'c_crossattn':343 [self.model.get_learned_conditioning([n_prompt] * num_samples)]344 }345 shape = (4, H // 8, W // 8)346 347 if config.save_memory:348 self.model.low_vram_shift(is_diffusing=True)349 350 samples, intermediates = self.ddim_sampler.sample(351 ddim_steps,352 num_samples,353 shape,354 cond,355 verbose=False,356 eta=eta,357 unconditional_guidance_scale=scale,358 unconditional_conditioning=un_cond)359 360 if config.save_memory:361 self.model.low_vram_shift(is_diffusing=False)362 363 x_samples = self.model.decode_first_stage(samples)364 x_samples = (365 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +366 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)367 368 results = [x_samples[i] for i in range(num_samples)]369 return [255 - detected_map] + results370 371 @torch.inference_mode()372 def process_scribble_interactive(self, input_image, prompt, a_prompt,373 n_prompt, num_samples, image_resolution,374 ddim_steps, scale, seed, eta):375 self.load_weight('scribble')376 377 img = resize_image(HWC3(input_image['mask'][:, :, 0]),378 image_resolution)379 H, W, C = img.shape380 381 detected_map = np.zeros_like(img, dtype=np.uint8)382 detected_map[np.min(img, axis=2) > 127] = 255383 384 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0385 control = torch.stack([control for _ in range(num_samples)], dim=0)386 control = einops.rearrange(control, 'b h w c -> b c h w').clone()387 388 if seed == -1:389 seed = random.randint(0, 65535)390 seed_everything(seed)391 392 if config.save_memory:393 self.model.low_vram_shift(is_diffusing=False)394 395 cond = {396 'c_concat': [control],397 'c_crossattn': [398 self.model.get_learned_conditioning(399 [prompt + ', ' + a_prompt] * num_samples)400 ]401 }402 un_cond = {403 'c_concat': [control],404 'c_crossattn':405 [self.model.get_learned_conditioning([n_prompt] * num_samples)]406 }407 shape = (4, H // 8, W // 8)408 409 if config.save_memory:410 self.model.low_vram_shift(is_diffusing=True)411 412 samples, intermediates = self.ddim_sampler.sample(413 ddim_steps,414 num_samples,415 shape,416 cond,417 verbose=False,418 eta=eta,419 unconditional_guidance_scale=scale,420 unconditional_conditioning=un_cond)421 422 if config.save_memory:423 self.model.low_vram_shift(is_diffusing=False)424 425 x_samples = self.model.decode_first_stage(samples)426 x_samples = (427 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +428 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)429 430 results = [x_samples[i] for i in range(num_samples)]431 return [255 - detected_map] + results432 433 @torch.inference_mode()434 def process_fake_scribble(self, input_image, prompt, a_prompt, n_prompt,435 num_samples, image_resolution, detect_resolution,436 ddim_steps, scale, seed, eta):437 self.load_weight('scribble')438 439 input_image = HWC3(input_image)440 detected_map = apply_hed(resize_image(input_image, detect_resolution))441 detected_map = HWC3(detected_map)442 img = resize_image(input_image, image_resolution)443 H, W, C = img.shape444 445 detected_map = cv2.resize(detected_map, (W, H),446 interpolation=cv2.INTER_LINEAR)447 detected_map = nms(detected_map, 127, 3.0)448 detected_map = cv2.GaussianBlur(detected_map, (0, 0), 3.0)449 detected_map[detected_map > 4] = 255450 detected_map[detected_map < 255] = 0451 452 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0453 control = torch.stack([control for _ in range(num_samples)], dim=0)454 control = einops.rearrange(control, 'b h w c -> b c h w').clone()455 456 if seed == -1:457 seed = random.randint(0, 65535)458 seed_everything(seed)459 460 if config.save_memory:461 self.model.low_vram_shift(is_diffusing=False)462 463 cond = {464 'c_concat': [control],465 'c_crossattn': [466 self.model.get_learned_conditioning(467 [prompt + ', ' + a_prompt] * num_samples)468 ]469 }470 un_cond = {471 'c_concat': [control],472 'c_crossattn':473 [self.model.get_learned_conditioning([n_prompt] * num_samples)]474 }475 shape = (4, H // 8, W // 8)476 477 if config.save_memory:478 self.model.low_vram_shift(is_diffusing=True)479 480 samples, intermediates = self.ddim_sampler.sample(481 ddim_steps,482 num_samples,483 shape,484 cond,485 verbose=False,486 eta=eta,487 unconditional_guidance_scale=scale,488 unconditional_conditioning=un_cond)489 490 if config.save_memory:491 self.model.low_vram_shift(is_diffusing=False)492 493 x_samples = self.model.decode_first_stage(samples)494 x_samples = (495 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +496 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)497 498 results = [x_samples[i] for i in range(num_samples)]499 return [255 - detected_map] + results500 501 @torch.inference_mode()502 def process_pose(self, input_image, prompt, a_prompt, n_prompt,503 num_samples, image_resolution, detect_resolution,504 ddim_steps, scale, seed, eta):505 self.load_weight('pose')506 507 input_image = HWC3(input_image)508 detected_map, _ = apply_openpose(509 resize_image(input_image, detect_resolution))510 detected_map = HWC3(detected_map)511 img = resize_image(input_image, image_resolution)512 H, W, C = img.shape513 514 detected_map = cv2.resize(detected_map, (W, H),515 interpolation=cv2.INTER_NEAREST)516 517 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0518 control = torch.stack([control for _ in range(num_samples)], dim=0)519 control = einops.rearrange(control, 'b h w c -> b c h w').clone()520 521 if seed == -1:522 seed = random.randint(0, 65535)523 seed_everything(seed)524 525 if config.save_memory:526 self.model.low_vram_shift(is_diffusing=False)527 528 cond = {529 'c_concat': [control],530 'c_crossattn': [531 self.model.get_learned_conditioning(532 [prompt + ', ' + a_prompt] * num_samples)533 ]534 }535 un_cond = {536 'c_concat': [control],537 'c_crossattn':538 [self.model.get_learned_conditioning([n_prompt] * num_samples)]539 }540 shape = (4, H // 8, W // 8)541 542 if config.save_memory:543 self.model.low_vram_shift(is_diffusing=True)544 545 samples, intermediates = self.ddim_sampler.sample(546 ddim_steps,547 num_samples,548 shape,549 cond,550 verbose=False,551 eta=eta,552 unconditional_guidance_scale=scale,553 unconditional_conditioning=un_cond)554 555 if config.save_memory:556 self.model.low_vram_shift(is_diffusing=False)557 558 x_samples = self.model.decode_first_stage(samples)559 x_samples = (560 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +561 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)562 563 results = [x_samples[i] for i in range(num_samples)]564 return [detected_map] + results565 566 @torch.inference_mode()567 def process_seg(self, input_image, prompt, a_prompt, n_prompt, num_samples,568 image_resolution, detect_resolution, ddim_steps, scale,569 seed, eta):570 self.load_weight('seg')571 572 input_image = HWC3(input_image)573 detected_map = apply_uniformer(574 resize_image(input_image, detect_resolution))575 img = resize_image(input_image, image_resolution)576 H, W, C = img.shape577 578 detected_map = cv2.resize(detected_map, (W, H),579 interpolation=cv2.INTER_NEAREST)580 581 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0582 control = torch.stack([control for _ in range(num_samples)], dim=0)583 control = einops.rearrange(control, 'b h w c -> b c h w').clone()584 585 if seed == -1:586 seed = random.randint(0, 65535)587 seed_everything(seed)588 589 if config.save_memory:590 self.model.low_vram_shift(is_diffusing=False)591 592 cond = {593 'c_concat': [control],594 'c_crossattn': [595 self.model.get_learned_conditioning(596 [prompt + ', ' + a_prompt] * num_samples)597 ]598 }599 un_cond = {600 'c_concat': [control],601 'c_crossattn':602 [self.model.get_learned_conditioning([n_prompt] * num_samples)]603 }604 shape = (4, H // 8, W // 8)605 606 if config.save_memory:607 self.model.low_vram_shift(is_diffusing=True)608 609 samples, intermediates = self.ddim_sampler.sample(610 ddim_steps,611 num_samples,612 shape,613 cond,614 verbose=False,615 eta=eta,616 unconditional_guidance_scale=scale,617 unconditional_conditioning=un_cond)618 619 if config.save_memory:620 self.model.low_vram_shift(is_diffusing=False)621 622 x_samples = self.model.decode_first_stage(samples)623 x_samples = (624 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +625 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)626 627 results = [x_samples[i] for i in range(num_samples)]628 return [detected_map] + results629 630 @torch.inference_mode()631 def process_depth(self, input_image, prompt, a_prompt, n_prompt,632 num_samples, image_resolution, detect_resolution,633 ddim_steps, scale, seed, eta):634 self.load_weight('depth')635 636 input_image = HWC3(input_image)637 detected_map, _ = apply_midas(638 resize_image(input_image, detect_resolution))639 detected_map = HWC3(detected_map)640 img = resize_image(input_image, image_resolution)641 H, W, C = img.shape642 643 detected_map = cv2.resize(detected_map, (W, H),644 interpolation=cv2.INTER_LINEAR)645 646 control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0647 control = torch.stack([control for _ in range(num_samples)], dim=0)648 control = einops.rearrange(control, 'b h w c -> b c h w').clone()649 650 if seed == -1:651 seed = random.randint(0, 65535)652 seed_everything(seed)653 654 if config.save_memory:655 self.model.low_vram_shift(is_diffusing=False)656 657 cond = {658 'c_concat': [control],659 'c_crossattn': [660 self.model.get_learned_conditioning(661 [prompt + ', ' + a_prompt] * num_samples)662 ]663 }664 un_cond = {665 'c_concat': [control],666 'c_crossattn':667 [self.model.get_learned_conditioning([n_prompt] * num_samples)]668 }669 shape = (4, H // 8, W // 8)670 671 if config.save_memory:672 self.model.low_vram_shift(is_diffusing=True)673 674 samples, intermediates = self.ddim_sampler.sample(675 ddim_steps,676 num_samples,677 shape,678 cond,679 verbose=False,680 eta=eta,681 unconditional_guidance_scale=scale,682 unconditional_conditioning=un_cond)683 684 if config.save_memory:685 self.model.low_vram_shift(is_diffusing=False)686 687 x_samples = self.model.decode_first_stage(samples)688 x_samples = (689 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +690 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)691 692 results = [x_samples[i] for i in range(num_samples)]693 return [detected_map] + results694 695 @torch.inference_mode()696 def process_normal(self, input_image, prompt, a_prompt, n_prompt,697 num_samples, image_resolution, detect_resolution,698 ddim_steps, scale, seed, eta, bg_threshold):699 self.load_weight('normal')700 701 input_image = HWC3(input_image)702 _, detected_map = apply_midas(resize_image(input_image,703 detect_resolution),704 bg_th=bg_threshold)705 detected_map = HWC3(detected_map)706 img = resize_image(input_image, image_resolution)707 H, W, C = img.shape708 709 detected_map = cv2.resize(detected_map, (W, H),710 interpolation=cv2.INTER_LINEAR)711 712 control = torch.from_numpy(713 detected_map[:, :, ::-1].copy()).float().cuda() / 255.0714 control = torch.stack([control for _ in range(num_samples)], dim=0)715 control = einops.rearrange(control, 'b h w c -> b c h w').clone()716 717 if seed == -1:718 seed = random.randint(0, 65535)719 seed_everything(seed)720 721 if config.save_memory:722 self.model.low_vram_shift(is_diffusing=False)723 724 cond = {725 'c_concat': [control],726 'c_crossattn': [727 self.model.get_learned_conditioning(728 [prompt + ', ' + a_prompt] * num_samples)729 ]730 }731 un_cond = {732 'c_concat': [control],733 'c_crossattn':734 [self.model.get_learned_conditioning([n_prompt] * num_samples)]735 }736 shape = (4, H // 8, W // 8)737 738 if config.save_memory:739 self.model.low_vram_shift(is_diffusing=True)740 741 samples, intermediates = self.ddim_sampler.sample(742 ddim_steps,743 num_samples,744 shape,745 cond,746 verbose=False,747 eta=eta,748 unconditional_guidance_scale=scale,749 unconditional_conditioning=un_cond)750 751 if config.save_memory:752 self.model.low_vram_shift(is_diffusing=False)753 754 x_samples = self.model.decode_first_stage(samples)755 x_samples = (756 einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +757 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)758 759 results = [x_samples[i] for i in range(num_samples)]760 return [detected_map] + results