fix: enhance the compatibility between diffusers and comfy

This commit is contained in:
Jianqi Pan
2024-07-31 22:12:08 +09:00
parent 4f6bb64067
commit 3c121b126c
8 changed files with 30 additions and 10 deletions
+19 -8
View File
@@ -201,7 +201,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -211,7 +211,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -221,7 +221,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -513,7 +513,7 @@ class DiffusersPrepareLatents:
batch_size=batch_size,
height=height,
width=width,
dtype=comfy.model_management.VAE_DTYPE,
dtype=comfy.model_management.vae_dtype(),
device=device,
generator=generator,
latents=latents,
@@ -688,8 +688,8 @@ class DiffusersControlNetUnitStack:
def run(
self,
controlnet_unit_1: tuple[ControlNetModel],
controlnet_unit_2: tuple[ControlNetModel] | None,
controlnet_unit_3: tuple[ControlNetModel] | None,
controlnet_unit_2: tuple[ControlNetModel] | None = None,
controlnet_unit_3: tuple[ControlNetModel] | None = None,
):
stack = []
if controlnet_unit_1:
@@ -749,13 +749,22 @@ class DiffusersGenerator:
"step": 64,
},
),
"reference_strength": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.01,
},
),
"reference_style_fidelity": (
"FLOAT",
{
"default": 0.5,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
},
),
},
@@ -801,6 +810,7 @@ class DiffusersGenerator:
reference_only_adain: str = "disable",
reference_image: torch.Tensor | None = None,
reference_style_fidelity: float = 0.5,
reference_strength: float = 1.0,
):
reference_only = reference_only == "enable"
reference_only_adain = reference_only_adain == "enable"
@@ -820,7 +830,7 @@ class DiffusersGenerator:
width=width,
generator=generator,
device=device,
dtype=comfy.model_management.VAE_DTYPE,
dtype=comfy.model_management.vae_dtype(),
)
images = latents_to_img_tensor(pipeline, latents)
else:
@@ -856,6 +866,7 @@ class DiffusersGenerator:
strength=strength,
controlnet_units=controlnet_units,
callback=callback,
reference_strength=reference_strength,
reference_attn=reference_only,
reference_adain=reference_only_adain,
style_fidelity=reference_style_fidelity,
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 362 KiB

After

Width:  |  Height:  |  Size: 416 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.4 MiB

After

Width:  |  Height:  |  Size: 2.4 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.1 MiB

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 792 KiB

After

Width:  |  Height:  |  Size: 806 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 687 KiB

After

Width:  |  Height:  |  Size: 685 KiB

+2 -2
View File
@@ -59,14 +59,14 @@ class PipelineWrapper:
self.pipeline = JannchiePipeline.from_single_file(
ckpt_path,
torch_dtype=unet_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers"),
cache_dir=folder_paths.get_folder_paths("diffusers")[0],
use_safetensors=True,
)
else:
self.pipeline = JannchiePipeline.from_pretrained(
ckpt_path,
torch_dtype=unet_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers"),
cache_dir=folder_paths.get_folder_paths("diffusers")[0],
use_safetensors=ckpt_path.endswith(".safetensors"),
)
+9
View File
@@ -227,6 +227,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
masked_image_latents: Optional[torch.FloatTensor] = None,
ip_adapter_image: Optional[PipelineImageInput] = None,
ip_adapter_image_embeds: Optional[List[torch.FloatTensor]] = None,
reference_strength: float = 1.0,
*arg,
**args,
):
@@ -515,6 +516,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
gn_auto_machine_weight=gn_auto_machine_weight,
ref_mask_dict=ref_mask_dict,
out_mask_dict=out_mask_dict,
strength=reference_strength,
)
if reference_attn:
self.unet = ReferenceOnlyUNet2DConditionModel.from_unet(
@@ -1132,6 +1134,7 @@ class ReferenceData:
gn_auto_machine_weight: float = 1.0
ref_mask_dict: dict = None
out_mask_dict: dict = None
strength: float = 1.0
class ReferenceOnlyUNet2DConditionModel(UNet2DConditionModel):
@@ -1310,6 +1313,7 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
style_fidelity * attn_output_c
+ (1.0 - style_fidelity) * attn_output_uc
)
attn_output *= ref_data.strength
bank.clear()
else:
# without reference only
@@ -1436,6 +1440,7 @@ class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
# apply additional residuals to the output of the last pair of resnet and attention blocks
if i == len(blocks) - 1 and additional_residuals is not None:
hidden_states = hidden_states + additional_residuals
@@ -1501,6 +1506,7 @@ class DownBlock2DReferenceOnly(DownBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
output_states = output_states + (hidden_states,)
@@ -1566,6 +1572,7 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
if do_classifier_free_guidance and style_fidelity > 0:
x_c[uc_mask] = x[uc_mask]
x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc
x *= self.ref_data.strength
self.mean_bank = []
self.var_bank = []
return x
@@ -1623,6 +1630,7 @@ class UpBlock2DReferenceOnly(UpBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
if MODE == "read":
self.mean_bank = []
@@ -1699,6 +1707,7 @@ class CrossAttnUpBlock2DReferenceOnly(CrossAttnUpBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
if MODE == "read":
self.mean_bank = []