From d3edc4a8f4c34284a524eba25ebba5008e497709 Mon Sep 17 00:00:00 2001 From: Jianqi Pan Date: Sun, 24 Mar 2024 15:35:14 +0900 Subject: [PATCH] :adhesive_bandage: fix(dtype): collect unet & vae type --- __init__.py | 4 +++- pipelines/__init__.py | 25 ++++++++++++++++++------- pipelines/jannchie.py | 13 +++++++++---- 3 files changed, 30 insertions(+), 12 deletions(-) diff --git a/__init__.py b/__init__.py index 862a935..9c3e004 100644 --- a/__init__.py +++ b/__init__.py @@ -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,) diff --git a/pipelines/__init__.py b/pipelines/__init__.py index bee7474..1baaff3 100644 --- a/pipelines/__init__.py +++ b/pipelines/__init__.py @@ -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() - diff --git a/pipelines/jannchie.py b/pipelines/jannchie.py index 96d2d32..0363eca 100644 --- a/pipelines/jannchie.py +++ b/pipelines/jannchie.py @@ -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 )