From b10a8540933fa56e1f8f065ea2e42b1f9ba9e839 Mon Sep 17 00:00:00 2001 From: Kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 25 Mar 2024 19:09:48 +0200 Subject: [PATCH] Add scheduler options --- nodes.py | 35 ++++++++++++++++++++++++++++++----- 1 file changed, 30 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 03178a1..961415d 100644 --- a/nodes.py +++ b/nodes.py @@ -4,7 +4,8 @@ import torch from omegaconf import OmegaConf from .models.depth_normal_pipeline_clip import DepthNormalEstimationPipeline from .models.unet_2d_condition import UNet2DConditionModel -from diffusers import DDIMScheduler, AutoencoderKL +from diffusers import DDIMScheduler, DDPMScheduler, DEISMultistepScheduler, PNDMScheduler, AutoencoderKL + import torch.nn.functional as F from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection @@ -25,7 +26,19 @@ def convert_dtype(dtype_str): return torch.bfloat16 else: raise NotImplementedError - + +scheduler_mapping = { + 'DDIMScheduler': DDIMScheduler, + 'DDPMScheduler': DDPMScheduler, + 'DEISMultistepScheduler': DEISMultistepScheduler, + 'PNDMScheduler': PNDMScheduler, + } +def get_scheduler_class(scheduler_str, model_path, subfolder='scheduler'): + if scheduler_str in scheduler_mapping: + return scheduler_mapping[scheduler_str].from_pretrained(model_path, subfolder=subfolder) + else: + raise ValueError(f"Unsupported scheduler: {scheduler_str}") + class geowizard_model_loader: @classmethod def INPUT_TYPES(s): @@ -38,6 +51,7 @@ class geowizard_model_loader: ], { "default": 'fp16' }), + }, } @@ -110,6 +124,15 @@ class geowizard_sampler: ], { "default": 'indoor' }), + "scheduler": ( + [ + 'DDIMScheduler', + 'DDPMScheduler', + 'DEISMultistepScheduler', + 'PNDMScheduler', + ], { + "default": 'DDIMScheduler' + }), "keep_model_loaded": ("BOOLEAN", {"default": True}), }, @@ -120,7 +143,7 @@ class geowizard_sampler: FUNCTION = "process" CATEGORY = "champWrapper" - def process(self, geowizard_model, image, domain, ensemble_size, steps, seed, keep_model_loaded): + def process(self, geowizard_model, image, domain, ensemble_size, steps, seed, scheduler, keep_model_loaded): device = mm.get_torch_device() mm.unload_all_models() mm.soft_empty_cache() @@ -128,7 +151,8 @@ class geowizard_sampler: dtype = pipe.dtype torch.manual_seed(seed) - + model_path = os.path.join(folder_paths.models_dir,'diffusers', 'geowizard') + pipe.scheduler = get_scheduler_class(scheduler, model_path) autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device) with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): @@ -178,10 +202,11 @@ class geowizard_sampler: depth_out = depth_out.cpu() normal_out = torch.cat(normal_maps, dim=0).cpu() + if not keep_model_loaded: + pipe = pipe.to('cpu') return (depth_out, normal_out) - NODE_CLASS_MAPPINGS = { "geowizard_model_loader": geowizard_model_loader, "geowizard_sampler": geowizard_sampler,