Team Ai
Apppublic

lee-t/ControlNet-Video

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
model.py760 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 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