From 91006653ea25adc8f0b3b91598f4d30b8e5054eb Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Dec 2023 15:08:22 +0200 Subject: [PATCH] Add scheduler options --- marigold/model/marigold_pipeline.py | 9 +++++++-- nodes.py | 15 ++++++++++++--- 2 files changed, 19 insertions(+), 5 deletions(-) diff --git a/marigold/model/marigold_pipeline.py b/marigold/model/marigold_pipeline.py index e010477..6a69b42 100644 --- a/marigold/model/marigold_pipeline.py +++ b/marigold/model/marigold_pipeline.py @@ -10,6 +10,7 @@ from diffusers import ( DDIMScheduler, DDPMScheduler, PNDMScheduler, + DEISMultistepScheduler, SchedulerMixin, UNet2DConditionModel, ) @@ -40,7 +41,7 @@ class MarigoldPipeline(nn.Module): trainable_unet=False, rgb_latent_scale_factor=0.18215, depth_latent_scale_factor=0.18215, - noise_scheduler_type="DDIMScheduler", + noise_scheduler_type=None, enable_gradient_checkpointing=False, enable_xformers=True, ) -> None: @@ -118,6 +119,11 @@ class MarigoldPipeline(nn.Module): noise_scheduler_pretrained_path["path"], subfolder=noise_scheduler_pretrained_path["subfolder"], ) + elif "DEISMultistepScheduler" == noise_scheduler_type: + self.noise_scheduler: SchedulerMixin = DEISMultistepScheduler.from_pretrained( + noise_scheduler_pretrained_path["path"], + subfolder=noise_scheduler_pretrained_path["subfolder"], + ) else: raise NotImplementedError @@ -251,7 +257,6 @@ class MarigoldPipeline(nn.Module): noise_pred = self.unet( unet_input, t, encoder_hidden_states=batch_empty_text_embed ).sample # [B, 4, h, w] - # compute the previous noisy sample x_t -> x_t-1 depth_latent = self.noise_scheduler.step( noise_pred, t, depth_latent diff --git a/nodes.py b/nodes.py index c63b7fb..60897cc 100644 --- a/nodes.py +++ b/nodes.py @@ -47,6 +47,15 @@ class MarigoldDepthEstimation: "keep_model_loaded": ("BOOLEAN", {"default": True}), "n_repeat_batch_size": ("INT", {"default": 2, "min": 1, "max": 4096, "step": 1}), "use_fp16": ("BOOLEAN", {"default": True}), + "scheduler": ( + [ + 'DDIMScheduler', + 'DDPMScheduler', + 'PNDMScheduler', + 'DEISMultistepScheduler', + ], { + "default": 'DDIMScheduler' + }), }, } @@ -57,7 +66,7 @@ class MarigoldDepthEstimation: CATEGORY = "Marigold" - def process(self, image, seed, denoise_steps, n_repeat, regularizer_strength, reduction_method, max_iter, tol,invert, keep_model_loaded, n_repeat_batch_size, use_fp16): + def process(self, image, seed, denoise_steps, n_repeat, regularizer_strength, reduction_method, max_iter, tol,invert, keep_model_loaded, n_repeat_batch_size, use_fp16, scheduler): batch_size = image.shape[0] precision = torch.float16 if use_fp16 else torch.float32 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -73,7 +82,7 @@ class MarigoldDepthEstimation: "../../models/diffusers/Marigold", ] - if not hasattr(self, 'marigold_pipeline') or self.marigold_pipeline is None or self.marigold_pipeline.unet.dtype != precision: + if not hasattr(self, 'marigold_pipeline') or self.marigold_pipeline is None or self.marigold_pipeline.unet.dtype != precision or self.marigold_pipeline.noise_scheduler != scheduler: # Load the model only if it hasn't been loaded before checkpoint_path = None for folder in folders_to_check: @@ -90,7 +99,7 @@ class MarigoldDepthEstimation: except: raise FileNotFoundError("No checkpoint directory found.") - self.marigold_pipeline = MarigoldPipeline.from_pretrained(checkpoint_path, enable_xformers=False, empty_text_embed=empty_text_embed) + self.marigold_pipeline = MarigoldPipeline.from_pretrained(checkpoint_path, enable_xformers=False, empty_text_embed=empty_text_embed, noise_scheduler_type=scheduler) self.marigold_pipeline = self.marigold_pipeline.to(device).half() if use_fp16 else self.marigold_pipeline.to(device) self.marigold_pipeline.unet.eval() # Set the model to evaluation mode pbar = comfy.utils.ProgressBar(batch_size * n_repeat)