Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions invokeai/app/invocations/denoise_latents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down Expand Up @@ -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),
Expand Down
31 changes: 22 additions & 9 deletions invokeai/backend/hidiffusion/hidiffusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,)
Expand Down Expand Up @@ -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
Expand Down Expand 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.
Expand All @@ -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
)
Expand Down Expand Up @@ -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)

Expand All @@ -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,
Expand Down
6 changes: 6 additions & 0 deletions invokeai/backend/stable_diffusion/extensions/hidiffusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down
4 changes: 4 additions & 0 deletions invokeai/backend/stable_diffusion/hidiffusion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
53 changes: 50 additions & 3 deletions tests/backend/stable_diffusion/test_hidiffusion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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": [],
Expand Down
Loading