🩹 fix(dtype): collect unet & vae type

This commit is contained in:
Jianqi Pan
2024-03-24 15:35:14 +09:00
parent 5be339b5a0
commit d3edc4a8f4
3 changed files with 30 additions and 12 deletions
+3 -1
View File
@@ -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
View File
@@ -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()
+9 -4
View File
@@ -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
)