diff --git a/__init__.py b/__init__.py index 4e0f0dd..2f950c5 100644 --- a/__init__.py +++ b/__init__.py @@ -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, diff --git a/examples/base.png b/examples/base.png index e26cb86..04a6b0f 100644 Binary files a/examples/base.png and b/examples/base.png differ diff --git a/examples/change_clothes.png b/examples/change_clothes.png index e9f1e47..b041bde 100644 Binary files a/examples/change_clothes.png and b/examples/change_clothes.png differ diff --git a/examples/controlnet.png b/examples/controlnet.png index 2eef4cf..f3803db 100644 Binary files a/examples/controlnet.png and b/examples/controlnet.png differ diff --git a/examples/inpainting.png b/examples/inpainting.png index 9a789c7..315f5fa 100644 Binary files a/examples/inpainting.png and b/examples/inpainting.png differ diff --git a/examples/reference_only.png b/examples/reference_only.png index bc61560..9d75459 100644 Binary files a/examples/reference_only.png and b/examples/reference_only.png differ diff --git a/pipelines/__init__.py b/pipelines/__init__.py index 35ea38d..74412a2 100644 --- a/pipelines/__init__.py +++ b/pipelines/__init__.py @@ -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"), ) diff --git a/pipelines/jannchie.py b/pipelines/jannchie.py index 0e1184a..335b6c6 100644 --- a/pipelines/jannchie.py +++ b/pipelines/jannchie.py @@ -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 = []