🩹 fix(dtype): collect unet & vae type
This commit is contained in:
+3
-1
@@ -74,6 +74,8 @@ def latents_to_img_tensor(pipeline, latents):
|
||||
# 1. 输入的 latents 是一个 -1 ~ 1 之间的 tensor
|
||||
# 2. 先进行缩放
|
||||
scaled_latents = latents / pipeline.vae.config.scaling_factor
|
||||
# 转成 unet 类型
|
||||
scaled_latents = scaled_latents.to(dtype=comfy.model_management.unet_dtype())
|
||||
# 3. 解码,返回的是 -1 ~ 1 之间的 tensor
|
||||
dec_tensor = pipeline.vae.decode(scaled_latents, return_dict=False)[0]
|
||||
# 4. 缩放到 0 ~ 1 之间
|
||||
@@ -509,7 +511,7 @@ class DiffusersControlNetLoader:
|
||||
cache_dir=folder_paths.get_folder_paths("controlnet")[0],
|
||||
).to(
|
||||
device=comfy.model_management.get_torch_device(),
|
||||
dtype=comfy.model_management.VAE_DTYPE,
|
||||
dtype=comfy.model_management.unet_dtype(),
|
||||
)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
+18
-7
@@ -1,6 +1,5 @@
|
||||
import contextlib
|
||||
|
||||
import comfy.model_management
|
||||
from diffusers import AutoencoderKL
|
||||
from diffusers.schedulers import (
|
||||
DEISMultistepScheduler,
|
||||
@@ -15,6 +14,9 @@ from diffusers.schedulers import (
|
||||
UniPCMultistepScheduler,
|
||||
)
|
||||
|
||||
import comfy.model_management
|
||||
import folder_paths
|
||||
|
||||
from .jannchie import *
|
||||
|
||||
schedulers = {
|
||||
@@ -47,27 +49,37 @@ class PipelineWrapper:
|
||||
):
|
||||
scheduler = schedulers.get(scheduler_name)
|
||||
device = comfy.model_management.get_torch_device()
|
||||
dtype = comfy.model_management.VAE_DTYPE
|
||||
vae_dtype = comfy.model_management.VAE_DTYPE
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
if ckpt_path.endswith(".safetensors"):
|
||||
self.pipeline = JannchiePipeline.from_single_file(
|
||||
ckpt_path,
|
||||
torch_dtype=dtype,
|
||||
torch_dtype=unet_dtype,
|
||||
cache_dir=folder_paths.get_folder_paths("diffusers"),
|
||||
use_safetensors=True,
|
||||
)
|
||||
else:
|
||||
self.pipeline = JannchiePipeline.from_pretrained(
|
||||
ckpt_path,
|
||||
torch_dtype=dtype,
|
||||
torch_dtype=unet_dtype,
|
||||
cache_dir=folder_paths.get_folder_paths("diffusers"),
|
||||
use_safetensors=ckpt_path.endswith(".safetensors"),
|
||||
)
|
||||
|
||||
if vae_path:
|
||||
if vae_path.endswith(".safetensors"):
|
||||
self.pipeline.vae = AutoencoderKL.from_single_file(
|
||||
vae_path,
|
||||
torch_dtype=dtype,
|
||||
torch_dtype=vae_dtype,
|
||||
cache_dir=folder_paths.get_folder_paths("diffusers"),
|
||||
use_safetensors=True,
|
||||
)
|
||||
else:
|
||||
self.pipeline.vae = AutoencoderKL.from_pretrained(
|
||||
vae_path,
|
||||
torch_dtype=dtype,
|
||||
torch_dtype=vae_dtype,
|
||||
cache_dir=folder_paths.get_folder_paths("diffusers"),
|
||||
use_safetensors=vae_path.endswith(".safetensors"),
|
||||
)
|
||||
if scheduler:
|
||||
self.pipeline.scheduler = scheduler
|
||||
@@ -75,4 +87,3 @@ class PipelineWrapper:
|
||||
self.pipeline.safety_checker = None
|
||||
with contextlib.suppress(Exception):
|
||||
self.pipeline.enable_xformers_memory_efficient_attention()
|
||||
|
||||
|
||||
@@ -228,7 +228,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
*arg,
|
||||
**args,
|
||||
):
|
||||
device = self._execution_device
|
||||
device = self.unet.device
|
||||
if height == None:
|
||||
if isinstance(image, torch.Tensor):
|
||||
if image is not None:
|
||||
@@ -441,14 +441,12 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
|
||||
# 7. Prepare mask latent variables
|
||||
if mask_image is not None:
|
||||
print(height, width)
|
||||
mask_condition = self.mask_processor.preprocess(
|
||||
mask_image, height=height, width=width
|
||||
)
|
||||
init_image = image
|
||||
init_image = init_image.to(dtype=torch.float32)
|
||||
if masked_image_latents is None:
|
||||
print(init_image.shape, mask_condition.shape)
|
||||
masked_image = init_image * (mask_condition < 0.5)
|
||||
else:
|
||||
masked_image = masked_image_latents
|
||||
@@ -458,7 +456,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
batch_size * num_images_per_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
self.unet.dtype,
|
||||
device,
|
||||
generator,
|
||||
do_classifier_free_guidance,
|
||||
@@ -666,6 +664,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
init_latents_proper = self.scheduler.add_noise(
|
||||
init_latents_proper, noise, torch.tensor([noise_timestep])
|
||||
)
|
||||
init_latents_proper = init_latents_proper.to(
|
||||
device=device, dtype=self.unet.dtype
|
||||
)
|
||||
|
||||
input_latents = (
|
||||
1 - init_mask
|
||||
@@ -686,6 +687,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if output_type != "latent":
|
||||
input_latents = input_latents.to(device=device, dtype=self.vae.dtype)
|
||||
result_imgs = self.vae.decode(
|
||||
input_latents / self.vae.config.scaling_factor, return_dict=False
|
||||
)[0]
|
||||
@@ -868,6 +870,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
]
|
||||
image_latents = torch.cat(image_latents, dim=0)
|
||||
else:
|
||||
image = image.to(self.vae.dtype)
|
||||
image_latents = self.vae.encode(image).latent_dist.sample(
|
||||
generator=generator
|
||||
)
|
||||
@@ -989,6 +992,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
image_latents = image_latents.repeat(
|
||||
batch_size // image_latents.shape[0], 1, 1, 1
|
||||
)
|
||||
image_latents.to(device=device, dtype=dtype)
|
||||
|
||||
if latents is None:
|
||||
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
@@ -1082,6 +1086,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
|
||||
]
|
||||
image_latents = torch.cat(image_latents, dim=0)
|
||||
else:
|
||||
image = image.to(self.vae.dtype)
|
||||
image_latents = retrieve_latents(
|
||||
self.vae.encode(image), generator=generator
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user