diff --git a/invokeai/app/invocations/denoise_latents.py b/invokeai/app/invocations/denoise_latents.py index 68e8dfc0dcd..413c7bc5fa4 100644 --- a/invokeai/app/invocations/denoise_latents.py +++ b/invokeai/app/invocations/denoise_latents.py @@ -920,6 +920,8 @@ def step_callback(state: PipelineIntermediateState) -> None: name_or_path=hidiffusion_name_or_path, apply_raunet=self.hidiffusion_raunet, apply_window_attn=self.hidiffusion_window_attn, + has_controlnet=bool(self.control), + is_controlnet_text_to_image=bool(self.control) and self.latents is None, t1_ratio=self.hidiffusion_t1_ratio, t2_ratio=self.hidiffusion_t2_ratio, generator=torch.Generator(device="cpu").manual_seed(seed), @@ -1146,6 +1148,8 @@ def _lora_loader() -> Iterator[PatchSpec]: name_or_path=hidiffusion_name_or_path, apply_raunet=self.hidiffusion_raunet, apply_window_attn=self.hidiffusion_window_attn, + has_controlnet=bool(controlnet_data), + is_controlnet_text_to_image=bool(controlnet_data) and self.latents is None, t1_ratio=self.hidiffusion_t1_ratio, t2_ratio=self.hidiffusion_t2_ratio, generator=torch.Generator(device="cpu").manual_seed(seed), diff --git a/invokeai/backend/hidiffusion/hidiffusion.py b/invokeai/backend/hidiffusion/hidiffusion.py index 5d67f2554ec..a5a53197111 100644 --- a/invokeai/backend/hidiffusion/hidiffusion.py +++ b/invokeai/backend/hidiffusion/hidiffusion.py @@ -897,6 +897,11 @@ def __call__( return sdxl_controlnet_ppl +def _resize_controlnet_residual(residual: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + """Resize a ControlNet residual to the current HiDiffusion feature-map size.""" + return F.interpolate(residual, target.shape[-2:], mode="bicubic") + + def make_diffusers_unet_2d_condition(block_class): class unet_2d_condition(block_class): # Save for unpatching later @@ -1203,9 +1208,8 @@ def forward( for down_block_res_sample, down_block_additional_residual in zip( down_block_res_samples, down_block_additional_residuals, strict=False ): - _, _, ori_H, ori_W = down_block_res_sample.shape - down_block_additional_residual = F.interpolate( - down_block_additional_residual, (ori_H, ori_W), mode="bicubic" + down_block_additional_residual = _resize_controlnet_residual( + down_block_additional_residual, down_block_res_sample ) down_block_res_sample = down_block_res_sample + down_block_additional_residual new_down_block_res_samples = new_down_block_res_samples + (down_block_res_sample,) @@ -1235,10 +1239,7 @@ def forward( sample += down_intrablock_additional_residuals.pop(0) if is_controlnet: - _, _, ori_H, ori_W = sample.shape - mid_block_additional_residual = F.interpolate( - mid_block_additional_residual, (ori_H, ori_W), mode="bicubic" - ) + mid_block_additional_residual = _resize_controlnet_residual(mid_block_additional_residual, sample) sample = sample + mid_block_additional_residual # 5. up @@ -2035,6 +2036,8 @@ def apply_hidiffusion( apply_window_attn: bool = True, is_playground=False, generator: torch.Generator | None = None, + has_controlnet: bool = False, + is_controlnet_text_to_image: bool = False, ): """ model: diffusers model. We support SD 1.5, 2.1, XL, XL Turbo. @@ -2052,7 +2055,9 @@ def apply_hidiffusion( if not is_diffusers: raise RuntimeError("Provided model was not a diffusers model/pipeline, as expected.") else: - # Check if the pipeline is a ControlNet pipeline + # Check if the pipeline is a ControlNet pipeline. InvokeAI's modular + # denoise passes a bare UNet, so it reports ControlNet separately. + has_controlnet = has_controlnet or hasattr(model, "controlnet") is_sdxl_controlnet = hasattr(model, "controlnet") and isinstance_str( model, "StableDiffusionXLControlNet", prefix=True ) @@ -2084,6 +2089,14 @@ def apply_hidiffusion( diffusion_model = model.unet if hasattr(model, "unet") else model + if ( + has_controlnet + and apply_raunet + and not (is_sdxl_controlnet_inpaint or is_sd_controlnet_inpaint or is_sdxl_controlnet or is_sd_controlnet) + ): + make_block_fn = make_diffusers_unet_2d_condition + diffusion_model.__class__ = make_block_fn(diffusion_model.__class__) + for _, module in diffusion_model.named_modules(): _snapshot_hidiffusion_state(module) @@ -2104,7 +2117,7 @@ def apply_hidiffusion( "size": None, "upsample_size": None, "hooks": [], - "text_to_img_controlnet": hasattr(model, "controlnet"), + "text_to_img_controlnet": has_controlnet and is_controlnet_text_to_image, "is_inpainting_task": model.__class__ in auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values(), "is_playground": is_playground, "pipeline": model, diff --git a/invokeai/backend/stable_diffusion/extensions/hidiffusion.py b/invokeai/backend/stable_diffusion/extensions/hidiffusion.py index 444c90f5480..13a1763f35e 100644 --- a/invokeai/backend/stable_diffusion/extensions/hidiffusion.py +++ b/invokeai/backend/stable_diffusion/extensions/hidiffusion.py @@ -20,11 +20,15 @@ def __init__( t1_ratio: Optional[float] = None, t2_ratio: Optional[float] = None, generator: torch.Generator | None = None, + has_controlnet: bool = False, + is_controlnet_text_to_image: bool = False, ): super().__init__() self._name_or_path = name_or_path self._apply_raunet = apply_raunet self._apply_window_attn = apply_window_attn + self._has_controlnet = has_controlnet + self._is_controlnet_text_to_image = is_controlnet_text_to_image self._t1_ratio = t1_ratio self._t2_ratio = t2_ratio self._generator = generator @@ -36,6 +40,8 @@ def patch_unet(self, unet: UNet2DConditionModel, original_weights: OriginalWeigh name_or_path=self._name_or_path, apply_raunet=self._apply_raunet, apply_window_attn=self._apply_window_attn, + has_controlnet=self._has_controlnet, + is_controlnet_text_to_image=self._is_controlnet_text_to_image, t1_ratio=self._t1_ratio, t2_ratio=self._t2_ratio, generator=self._generator, diff --git a/invokeai/backend/stable_diffusion/hidiffusion_utils.py b/invokeai/backend/stable_diffusion/hidiffusion_utils.py index 327d7d083b1..f6e6e1681b3 100644 --- a/invokeai/backend/stable_diffusion/hidiffusion_utils.py +++ b/invokeai/backend/stable_diffusion/hidiffusion_utils.py @@ -24,6 +24,8 @@ def hidiffusion_patch( t1_ratio: Optional[float] = None, t2_ratio: Optional[float] = None, generator: torch.Generator | None = None, + has_controlnet: bool = False, + is_controlnet_text_to_image: bool = False, ): """Context manager that applies HiDiffusion and restores the model on exit.""" from invokeai.backend.hidiffusion.hidiffusion import apply_hidiffusion, remove_hidiffusion @@ -118,6 +120,8 @@ def _apply_ratio_overrides(ratio_dict: dict) -> None: model, apply_raunet=apply_raunet, apply_window_attn=apply_window_attn, + has_controlnet=has_controlnet, + is_controlnet_text_to_image=is_controlnet_text_to_image, generator=generator, ) yield diff --git a/tests/backend/stable_diffusion/test_hidiffusion_utils.py b/tests/backend/stable_diffusion/test_hidiffusion_utils.py index 74c92ab9604..5f8619a4882 100644 --- a/tests/backend/stable_diffusion/test_hidiffusion_utils.py +++ b/tests/backend/stable_diffusion/test_hidiffusion_utils.py @@ -6,12 +6,13 @@ import torch from invokeai.backend.hidiffusion.hidiffusion import ( - remove_hidiffusion as real_remove_hidiffusion, -) -from invokeai.backend.hidiffusion.hidiffusion import ( + _resize_controlnet_residual, switching_threshold_ratio_dict, text_to_img_controlnet_switching_threshold_ratio_dict, ) +from invokeai.backend.hidiffusion.hidiffusion import ( + remove_hidiffusion as real_remove_hidiffusion, +) from invokeai.backend.stable_diffusion.hidiffusion_utils import hidiffusion_patch @@ -124,6 +125,52 @@ def run_with_global_seed(global_seed: int) -> torch.Tensor: torch.testing.assert_close(first, second) +@pytest.mark.parametrize("is_text_to_image", [False, True]) +def test_hidiffusion_patch_uses_controlnet_aware_forward_for_bare_unet(is_text_to_image: bool): + model = ModelMixin() + original_class = model.__class__ + + with hidiffusion_patch( + model, + name_or_path="stabilityai/stable-diffusion-xl-base-1.0", + apply_raunet=True, + apply_window_attn=False, + has_controlnet=True, + is_controlnet_text_to_image=is_text_to_image, + ): + assert model.__class__ is not original_class + assert model.__class__._parent is original_class + assert model.info["text_to_img_controlnet"] is is_text_to_image + + assert model.__class__ is original_class + + +def test_hidiffusion_patch_does_not_replace_unet_forward_for_window_attention_only(): + model = ModelMixin() + original_class = model.__class__ + + with hidiffusion_patch( + model, + name_or_path="stabilityai/stable-diffusion-xl-base-1.0", + apply_raunet=False, + apply_window_attn=True, + has_controlnet=True, + ): + assert model.__class__ is original_class + + +@pytest.mark.parametrize(("source_size", "target_size"), [(32, 16), (46, 23), (48, 24), (64, 32)]) +def test_hidiffusion_resizes_controlnet_residuals_to_current_feature_map(source_size: int, target_size: int): + residual = torch.ones(1, 2, source_size, source_size) + feature_map = torch.zeros(1, 2, target_size, target_size) + + resized_residual = _resize_controlnet_residual(residual, feature_map) + combined = feature_map + resized_residual + + assert combined.shape == feature_map.shape + torch.testing.assert_close(combined, torch.ones_like(feature_map)) + + def test_hidiffusion_patch_resets_cached_runtime_state_when_reenabled(): module_keys = { "down_module_key": [],