Add scheduler options

This commit is contained in:
kijai
2023-12-27 15:08:22 +02:00
parent d98ba12708
commit 91006653ea
2 changed files with 19 additions and 5 deletions
+7 -2
View File
@@ -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
+12 -3
View File
@@ -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)