Compatible with the latest version of ComfyUI
This commit is contained in:
+3
-5
@@ -40,14 +40,12 @@ def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive,
|
||||
else:
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent['noise_mask']
|
||||
noise_mask = torch.nn.functional.interpolate(noise_mask[None, None,], size=(noise.shape[2], noise.shape[3]),
|
||||
mode="bilinear")
|
||||
noise_mask = torch.nn.functional.interpolate(noise_mask[None, None,], size=(noise.shape[2], noise.shape[3]), mode="bilinear")
|
||||
noise_mask = noise_mask.round()
|
||||
noise_mask = torch.cat([noise_mask] * noise.shape[1], dim=1)
|
||||
noise_mask = torch.cat([noise_mask] * noise.shape[0])
|
||||
noise_mask = noise_mask.to(device)
|
||||
|
||||
|
||||
real_model = None
|
||||
model_management.load_model_gpu(model)
|
||||
real_model = model.model
|
||||
@@ -82,7 +80,7 @@ def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive,
|
||||
model_management.load_controlnet_gpu(control_net_models)
|
||||
|
||||
if sampler_name in comfy.samplers.KSampler.SAMPLERS:
|
||||
sampler = comfy.samplers.KSampler(real_model, steps=steps, device=device, sampler=sampler_name, scheduler=scheduler, denoise=denoise)
|
||||
sampler = comfy.samplers.KSampler(real_model, steps=steps, device=device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options)
|
||||
else:
|
||||
#other samplers
|
||||
pass
|
||||
@@ -254,7 +252,7 @@ class CheckpointLoaderSimpleSequence:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
|
||||
RETURN_TYPES = ("ORIGINAL_MODEL", "CLIP", "VAE")
|
||||
FUNCTION = "load_checkpoint"
|
||||
|
||||
CATEGORY = "vid2vid"
|
||||
|
||||
@@ -140,12 +140,12 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, o
|
||||
model = instantiate_from_config(model_config)
|
||||
model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to)
|
||||
|
||||
#with torch.inference_mode(mode=False):
|
||||
with torch.inference_mode(mode=False):
|
||||
model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config}))
|
||||
if model_management.xformers_enabled():
|
||||
model.model.diffusion_model.enable_xformers_memory_efficient_attention()
|
||||
|
||||
if fp16:
|
||||
model = model.half()
|
||||
#if fp16:
|
||||
# model = model.half()
|
||||
|
||||
return (ModelPatcher(model), clip, vae, clipvision)
|
||||
|
||||
+1
-1
@@ -21,7 +21,7 @@ def next_step(model_output: Union[torch.FloatTensor, np.ndarray], timestep: int,
|
||||
|
||||
def get_noise_pred_single(latents, t, context, unet):
|
||||
latents = rearrange(latents.squeeze(0), "c f h w -> f c h w")
|
||||
noise_pred = unet(latents, t.view(1), context=context)
|
||||
noise_pred = unet(latents, t.view(1).to(latents.device), context=context)
|
||||
noise_pred = rearrange(noise_pred.unsqueeze(0), "b f c h w -> b c f h w")
|
||||
return noise_pred
|
||||
|
||||
|
||||
Reference in New Issue
Block a user