fix: enhance the compatibility between diffusers and comfy
@@ -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,
|
||||
|
||||
|
Before Width: | Height: | Size: 362 KiB After Width: | Height: | Size: 416 KiB |
|
Before Width: | Height: | Size: 2.4 MiB After Width: | Height: | Size: 2.4 MiB |
|
Before Width: | Height: | Size: 1.1 MiB After Width: | Height: | Size: 1.2 MiB |
|
Before Width: | Height: | Size: 792 KiB After Width: | Height: | Size: 806 KiB |
|
Before Width: | Height: | Size: 687 KiB After Width: | Height: | Size: 685 KiB |
@@ -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"),
|
||||
)
|
||||
|
||||
|
||||
@@ -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 = []
|
||||
|
||||