diff --git a/docs/src/content/docs/features/hidiffusion.mdx b/docs/src/content/docs/features/hidiffusion.mdx index 8158f4980a1..1a8a35b276b 100644 --- a/docs/src/content/docs/features/hidiffusion.mdx +++ b/docs/src/content/docs/features/hidiffusion.mdx @@ -17,6 +17,7 @@ Learn more: https://github.com/megvii-research/HiDiffusion 3. In the **Advanced** grid, enable **HiDiffusion** and optionally adjust the two sub‑toggles and ratios: - **HiDiffusion: RAU‑Net** - **HiDiffusion: Window Attention** + - **HiDiffusion: Automatic Ratios** - **HiDiffusion: T1 Ratio** - **HiDiffusion: T2 Ratio** @@ -28,9 +29,11 @@ Learn more: https://github.com/megvii-research/HiDiffusion - **HiDiffusion: Window Attention**: Enables windowed attention blocks. This can boost local texture/detail, but may slightly affect global coherence in some prompts. -- **HiDiffusion: T1 Ratio**: Controls when HiDiffusion switches into its mid‑stage behavior. Lower values switch earlier; higher values preserve global structure longer. +- **HiDiffusion: Automatic Ratios**: Uses HiDiffusion's model- and resolution-specific T1/T2 presets. Disable it to enter manual overrides. -- **HiDiffusion: T2 Ratio**: Controls when HiDiffusion switches into its late‑stage behavior. Higher values keep window attention active longer and can sharpen local detail. +- **HiDiffusion: T1 Ratio**: Controls the first RAU-Net switching threshold. Lower values switch earlier; higher values keep the first resolution-aware stage active longer. + +- **HiDiffusion: T2 Ratio**: Controls the second RAU-Net switching threshold used for extreme-resolution generation. It does not control window attention. Higher values keep the second resolution-aware stage active longer. ## Tips diff --git a/invokeai/app/invocations/denoise_latents.py b/invokeai/app/invocations/denoise_latents.py index 2d48dd87607..e00e48cbbb4 100644 --- a/invokeai/app/invocations/denoise_latents.py +++ b/invokeai/app/invocations/denoise_latents.py @@ -133,7 +133,7 @@ def get_scheduler( title="Denoise - SD1.5, SDXL", tags=["latents", "denoise", "txt2img", "t2i", "t2l", "img2img", "i2i", "l2l"], category="latents", - version="1.6.0", + version="1.7.0", ) class DenoiseLatentsInvocation(BaseInvocation): """Denoises noisy latents to decodable images""" @@ -209,15 +209,15 @@ class DenoiseLatentsInvocation(BaseInvocation): description=FieldDescriptions.hidiffusion_window_attn, title="HiDiffusion: Window Attention", ) - hidiffusion_t1_ratio: float = InputField( - default=0.4, + hidiffusion_t1_ratio: Optional[float] = InputField( + default=None, ge=0, le=1, description=FieldDescriptions.hidiffusion_t1_ratio, title="HiDiffusion: T1 Ratio", ) - hidiffusion_t2_ratio: float = InputField( - default=0.0, + hidiffusion_t2_ratio: Optional[float] = InputField( + default=None, ge=0, le=1, description=FieldDescriptions.hidiffusion_t2_ratio, @@ -926,6 +926,10 @@ def step_callback(state: PipelineIntermediateState) -> None: t1_ratio=self.hidiffusion_t1_ratio, t2_ratio=self.hidiffusion_t2_ratio, generator=torch.Generator(device="cpu").manual_seed(seed), + is_inpainting_task=self.denoise_mask is not None, + use_aggressive_raunet=False, + denoising_start=self.denoising_start, + denoising_end=self.denoising_end, ) ) @@ -1157,6 +1161,10 @@ def _lora_loader() -> Iterator[PatchSpec]: t1_ratio=self.hidiffusion_t1_ratio, t2_ratio=self.hidiffusion_t2_ratio, generator=torch.Generator(device="cpu").manual_seed(seed), + is_inpainting_task=self.denoise_mask is not None, + use_aggressive_raunet=False, + denoising_start=self.denoising_start, + denoising_end=self.denoising_end, ) if self.hidiffusion else nullcontext() diff --git a/invokeai/app/invocations/metadata_linked.py b/invokeai/app/invocations/metadata_linked.py index 5f8c18ecc98..be23ab6d85e 100644 --- a/invokeai/app/invocations/metadata_linked.py +++ b/invokeai/app/invocations/metadata_linked.py @@ -624,7 +624,7 @@ class LatentsMetaOutput(LatentsOutput, MetadataOutput): title=f"{DenoiseLatentsInvocation.UIConfig.title} + Metadata", tags=["latents", "denoise", "txt2img", "t2i", "t2l", "img2img", "i2i", "l2l"], category="metadata", - version="1.2.0", + version="1.3.0", ) class DenoiseLatentsMetaInvocation(DenoiseLatentsInvocation, WithMetadata): def invoke(self, context: InvocationContext) -> LatentsMetaOutput: diff --git a/invokeai/backend/hidiffusion/hidiffusion.py b/invokeai/backend/hidiffusion/hidiffusion.py index 33437853214..6bd8fec67b0 100644 --- a/invokeai/backend/hidiffusion/hidiffusion.py +++ b/invokeai/backend/hidiffusion/hidiffusion.py @@ -136,6 +136,32 @@ def _get_max_timesteps(info_dict: dict) -> int: return len(pipeline.scheduler.timesteps) +def _get_switching_threshold_ratio(module: torch.nn.Module, presets: dict, preset_key: str) -> float: + """Resolve a threshold ratio for the executed part of the denoising schedule.""" + override = module.info["switching_threshold_overrides"].get(module.switching_threshold_ratio) + full_schedule_ratio = override if override is not None else presets[preset_key][module.switching_threshold_ratio] + + denoising_start = module.info.get("denoising_start", 0.0) + denoising_end = module.info.get("denoising_end", 1.0) + if denoising_end <= denoising_start: + return 0.0 + + executed_schedule_ratio = (full_schedule_ratio - denoising_start) / (denoising_end - denoising_start) + return max(0.0, min(1.0, executed_schedule_ratio)) + + +def _should_use_aggressive_raunet(module: torch.nn.Module) -> bool: + """Resolve whether RAU-Net should be activated after denoising has already started.""" + override = module.info.get("use_aggressive_raunet") + if override is not None: + return override + if module.info["is_inpainting_task"]: + return inpainting_is_aggressive_raunet + if module.info["is_playground"]: + return playground_is_aggressive_raunet + return is_aggressive_raunet + + def make_diffusers_sdxl_controlnet_ppl(block_class): class sdxl_controlnet_ppl(block_class): # Save for unpatching later @@ -1559,32 +1585,29 @@ def forward( ori_H, ori_W = self.info["size"] if self.model == "sd15": if ori_H < 256 or ori_W < 256: - self.T1_ratio = switching_threshold_ratio_dict["sd15_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_1024") else: - self.T1_ratio = switching_threshold_ratio_dict["sd15_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_2048") elif self.model == "sdxl": if ori_H < 512 or ori_W < 512: if self.info["text_to_img_controlnet"]: - self.T1_ratio = text_to_img_controlnet_switching_threshold_ratio_dict["sdxl_2048"][ - self.switching_threshold_ratio - ] + self.T1_ratio = _get_switching_threshold_ratio( + self, text_to_img_controlnet_switching_threshold_ratio_dict, "sdxl_2048" + ) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio( + self, switching_threshold_ratio_dict, "sdxl_2048" + ) - if self.info["is_inpainting_task"]: - self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info["is_playground"]: - self.aggressive_raunet = playground_is_aggressive_raunet - else: - self.aggressive_raunet = is_aggressive_raunet + self.aggressive_raunet = _should_use_aggressive_raunet(self) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_4096"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_4096") elif self.model == "sdxl_turbo": - self.T1_ratio = switching_threshold_ratio_dict["sdxl_turbo_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_turbo_1024") else: raise Exception("Error model. HiDiffusion now only supports sd15, sd21, sdxl, sdxl-turbo.") - if self.aggressive_raunet: + if self.aggressive_raunet and self.switching_threshold_ratio == "T1_ratio": # self.T1_start = min(int(self.max_timestep * self.T1_ratio * 0.4), int(8/50 * self.max_timestep)) self.T1_start = int(aggressive_step / 50 * self.max_timestep) self.T1_end = int(self.max_timestep * self.T1_ratio) @@ -1693,33 +1716,30 @@ def forward( ori_H, ori_W = self.info["size"] if self.model == "sd15": if ori_H < 256 or ori_W < 256: - self.T1_ratio = switching_threshold_ratio_dict["sd15_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_1024") else: - self.T1_ratio = switching_threshold_ratio_dict["sd15_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_2048") elif self.model == "sdxl": if ori_H < 512 or ori_W < 512: if self.info["text_to_img_controlnet"]: - self.T1_ratio = text_to_img_controlnet_switching_threshold_ratio_dict["sdxl_2048"][ - self.switching_threshold_ratio - ] + self.T1_ratio = _get_switching_threshold_ratio( + self, text_to_img_controlnet_switching_threshold_ratio_dict, "sdxl_2048" + ) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio( + self, switching_threshold_ratio_dict, "sdxl_2048" + ) - if self.info["is_inpainting_task"]: - self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info["is_playground"]: - self.aggressive_raunet = playground_is_aggressive_raunet - else: - self.aggressive_raunet = is_aggressive_raunet + self.aggressive_raunet = _should_use_aggressive_raunet(self) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_4096"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_4096") elif self.model == "sdxl_turbo": - self.T1_ratio = switching_threshold_ratio_dict["sdxl_turbo_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_turbo_1024") else: raise Exception("Error model. HiDiffusion now only supports sd15, sd21, sdxl, sdxl-turbo.") - if self.aggressive_raunet: + if self.aggressive_raunet and self.switching_threshold_ratio == "T1_ratio": # self.T1_start = min(int(self.max_timestep * self.T1_ratio * 0.4), int(8/50 * self.max_timestep)) self.T1_start = int(aggressive_step / 50 * self.max_timestep) self.T1_end = int(self.max_timestep * self.T1_ratio) @@ -1830,32 +1850,29 @@ def forward(self, hidden_states: torch.Tensor, scale=1.0) -> torch.Tensor: ori_H, ori_W = self.info["size"] if self.model == "sd15": if ori_H < 256 or ori_W < 256: - self.T1_ratio = switching_threshold_ratio_dict["sd15_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_1024") else: - self.T1_ratio = switching_threshold_ratio_dict["sd15_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_2048") elif self.model == "sdxl": if ori_H < 512 or ori_W < 512: if self.info["text_to_img_controlnet"]: - self.T1_ratio = text_to_img_controlnet_switching_threshold_ratio_dict["sdxl_2048"][ - self.switching_threshold_ratio - ] + self.T1_ratio = _get_switching_threshold_ratio( + self, text_to_img_controlnet_switching_threshold_ratio_dict, "sdxl_2048" + ) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio( + self, switching_threshold_ratio_dict, "sdxl_2048" + ) - if self.info["is_inpainting_task"]: - self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info["is_playground"]: - self.aggressive_raunet = playground_is_aggressive_raunet - else: - self.aggressive_raunet = is_aggressive_raunet + self.aggressive_raunet = _should_use_aggressive_raunet(self) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_4096"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_4096") elif self.model == "sdxl_turbo": - self.T1_ratio = switching_threshold_ratio_dict["sdxl_turbo_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_turbo_1024") else: raise Exception("Error model. HiDiffusion now only supports sd15, sd21, sdxl, sdxl-turbo.") - if self.aggressive_raunet: + if self.aggressive_raunet and self.switching_threshold_ratio == "T1_ratio": # self.T1 = min(int(self.max_timestep * self.T1_ratio), int(8/50 * self.max_timestep)) self.T1 = int(aggressive_step / 50 * self.max_timestep) else: @@ -1911,32 +1928,29 @@ def forward(self, hidden_states: torch.Tensor, scale=1.0) -> torch.Tensor: ori_H, ori_W = self.info["size"] if self.model == "sd15": if ori_H < 256 or ori_W < 256: - self.T1_ratio = switching_threshold_ratio_dict["sd15_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_1024") else: - self.T1_ratio = switching_threshold_ratio_dict["sd15_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sd15_2048") elif self.model == "sdxl": if ori_H < 512 or ori_W < 512: if self.info["text_to_img_controlnet"]: - self.T1_ratio = text_to_img_controlnet_switching_threshold_ratio_dict["sdxl_2048"][ - self.switching_threshold_ratio - ] + self.T1_ratio = _get_switching_threshold_ratio( + self, text_to_img_controlnet_switching_threshold_ratio_dict, "sdxl_2048" + ) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_2048"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio( + self, switching_threshold_ratio_dict, "sdxl_2048" + ) - if self.info["is_inpainting_task"]: - self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info["is_playground"]: - self.aggressive_raunet = playground_is_aggressive_raunet - else: - self.aggressive_raunet = is_aggressive_raunet + self.aggressive_raunet = _should_use_aggressive_raunet(self) else: - self.T1_ratio = switching_threshold_ratio_dict["sdxl_4096"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_4096") elif self.model == "sdxl_turbo": - self.T1_ratio = switching_threshold_ratio_dict["sdxl_turbo_1024"][self.switching_threshold_ratio] + self.T1_ratio = _get_switching_threshold_ratio(self, switching_threshold_ratio_dict, "sdxl_turbo_1024") else: raise Exception("Error model. HiDiffusion now only supports sd15, sd21, sdxl, sdxl-turbo.") - if self.aggressive_raunet: + if self.aggressive_raunet and self.switching_threshold_ratio == "T1_ratio": # self.T1 = min(int(self.max_timestep * self.T1_ratio), int(8/50 * self.max_timestep)) self.T1 = int(aggressive_step / 50 * self.max_timestep) else: @@ -2045,6 +2059,12 @@ def apply_hidiffusion( generator: torch.Generator | None = None, has_controlnet: bool = False, is_controlnet_text_to_image: bool = False, + t1_ratio: float | None = None, + t2_ratio: float | None = None, + is_inpainting_task: bool | None = None, + use_aggressive_raunet: bool | None = None, + denoising_start: float = 0.0, + denoising_end: float = 1.0, ): """ model: diffusers model. We support SD 1.5, 2.1, XL, XL Turbo. @@ -2120,14 +2140,19 @@ def apply_hidiffusion( elif set(sdxl_module_key) < set(diffusion_model_module_key): name_or_path = "stabilityai/stable-diffusion-xl-base-1.0" + detected_inpainting_task = model.__class__ in auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values() diffusion_model.info = { "size": None, "upsample_size": None, "hooks": [], "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_inpainting_task": detected_inpainting_task if is_inpainting_task is None else is_inpainting_task, "is_playground": is_playground, + "use_aggressive_raunet": use_aggressive_raunet, + "denoising_start": denoising_start, + "denoising_end": denoising_end, "pipeline": model, + "switching_threshold_overrides": {"T1_ratio": t1_ratio, "T2_ratio": t2_ratio}, } model.info = diffusion_model.info hook_diffusion_model(diffusion_model) diff --git a/invokeai/backend/stable_diffusion/diffusion/regional_ip_data.py b/invokeai/backend/stable_diffusion/diffusion/regional_ip_data.py index 792c97114da..eb2a9eea56b 100644 --- a/invokeai/backend/stable_diffusion/diffusion/regional_ip_data.py +++ b/invokeai/backend/stable_diffusion/diffusion/regional_ip_data.py @@ -11,9 +11,12 @@ def __init__( masks: list[torch.Tensor], dtype: torch.dtype, device: torch.device, - max_downscale_factor: int = 8, + max_downscale_factor: int = 16, ): - """Initialize a `IPAdapterConditioningData` object.""" + """Initialize an `IPAdapterConditioningData` object. + + HiDiffusion's RAU-Net requires one mask level beyond the standard UNet's 8x downscale. + """ assert len(image_prompt_embeds) == len(scales) == len(masks) # The image prompt embeddings. diff --git a/invokeai/backend/stable_diffusion/diffusion/regional_prompt_data.py b/invokeai/backend/stable_diffusion/diffusion/regional_prompt_data.py index eddd31f0c42..bebe00b35b1 100644 --- a/invokeai/backend/stable_diffusion/diffusion/regional_prompt_data.py +++ b/invokeai/backend/stable_diffusion/diffusion/regional_prompt_data.py @@ -19,7 +19,7 @@ def __init__( regions: list[TextConditioningRegions], device: torch.device, dtype: torch.dtype, - max_downscale_factor: int = 8, + max_downscale_factor: int = 16, ): """Initialize a `RegionalPromptData` object. Args: @@ -28,7 +28,7 @@ def __init__( device (torch.device): The device to use for the attention masks. dtype (torch.dtype): The data type to use for the attention masks. max_downscale_factor: Spatial masks will be prepared for downscale factors from 1 to max_downscale_factor - in steps of 2x. + in steps of 2x. HiDiffusion's RAU-Net requires one level beyond the standard UNet's 8x downscale. """ self._regions = regions self._device = device @@ -41,7 +41,7 @@ def __init__( self._negative_cross_attn_mask_score = -10000.0 def _prepare_spatial_masks( - self, regions: list[TextConditioningRegions], max_downscale_factor: int = 8 + self, regions: list[TextConditioningRegions], max_downscale_factor: int = 16 ) -> list[dict[int, torch.Tensor]]: """Prepare the spatial masks for all downscaling factors.""" # batch_masks_by_seq_len[b][s] contains the spatial masks for the b'th batch sample with a query sequence length diff --git a/invokeai/backend/stable_diffusion/extensions/hidiffusion.py b/invokeai/backend/stable_diffusion/extensions/hidiffusion.py index 13a1763f35e..464e6a1d197 100644 --- a/invokeai/backend/stable_diffusion/extensions/hidiffusion.py +++ b/invokeai/backend/stable_diffusion/extensions/hidiffusion.py @@ -22,6 +22,10 @@ def __init__( generator: torch.Generator | None = None, has_controlnet: bool = False, is_controlnet_text_to_image: bool = False, + is_inpainting_task: bool | None = None, + use_aggressive_raunet: bool | None = None, + denoising_start: float = 0.0, + denoising_end: float = 1.0, ): super().__init__() self._name_or_path = name_or_path @@ -29,6 +33,10 @@ def __init__( self._apply_window_attn = apply_window_attn self._has_controlnet = has_controlnet self._is_controlnet_text_to_image = is_controlnet_text_to_image + self._is_inpainting_task = is_inpainting_task + self._use_aggressive_raunet = use_aggressive_raunet + self._denoising_start = denoising_start + self._denoising_end = denoising_end self._t1_ratio = t1_ratio self._t2_ratio = t2_ratio self._generator = generator @@ -45,5 +53,9 @@ def patch_unet(self, unet: UNet2DConditionModel, original_weights: OriginalWeigh t1_ratio=self._t1_ratio, t2_ratio=self._t2_ratio, generator=self._generator, + is_inpainting_task=self._is_inpainting_task, + use_aggressive_raunet=self._use_aggressive_raunet, + denoising_start=self._denoising_start, + denoising_end=self._denoising_end, ): yield None diff --git a/invokeai/backend/stable_diffusion/hidiffusion_utils.py b/invokeai/backend/stable_diffusion/hidiffusion_utils.py index f6e6e1681b3..371da5ca6bd 100644 --- a/invokeai/backend/stable_diffusion/hidiffusion_utils.py +++ b/invokeai/backend/stable_diffusion/hidiffusion_utils.py @@ -1,19 +1,11 @@ from __future__ import annotations -import copy import sys from contextlib import contextmanager from typing import Any, Optional import torch -from invokeai.backend.hidiffusion.hidiffusion import ( - switching_threshold_ratio_dict as _switching_threshold_ratio_dict, -) -from invokeai.backend.hidiffusion.hidiffusion import ( - text_to_img_controlnet_switching_threshold_ratio_dict as _text_to_img_controlnet_switching_threshold_ratio_dict, -) - @contextmanager def hidiffusion_patch( @@ -26,6 +18,10 @@ def hidiffusion_patch( generator: torch.Generator | None = None, has_controlnet: bool = False, is_controlnet_text_to_image: bool = False, + is_inpainting_task: bool | None = None, + use_aggressive_raunet: bool | None = None, + denoising_start: float = 0.0, + denoising_end: float = 1.0, ): """Context manager that applies HiDiffusion and restores the model on exit.""" from invokeai.backend.hidiffusion.hidiffusion import apply_hidiffusion, remove_hidiffusion @@ -70,14 +66,6 @@ def _set_name_or_path_on_config(cfg, value: str) -> bool: original_num_upsamplers = getattr(target, "num_upsamplers", None) - ratio_overrides = None - ratio_dicts = None - if t1_ratio is not None or t2_ratio is not None: - ratio_dicts = ( - _switching_threshold_ratio_dict, - _text_to_img_controlnet_switching_threshold_ratio_dict, - ) - set_model_name_or_path = False set_config_name_or_path = False try: @@ -103,26 +91,19 @@ def _set_name_or_path_on_config(cfg, value: str) -> bool: except Exception: pass - if ratio_dicts is not None: - ratio_overrides = (copy.deepcopy(ratio_dicts[0]), copy.deepcopy(ratio_dicts[1])) - - def _apply_ratio_overrides(ratio_dict: dict) -> None: - for _, entry in ratio_dict.items(): - if t1_ratio is not None: - entry["T1_ratio"] = t1_ratio - if t2_ratio is not None and "T2_ratio" in entry: - entry["T2_ratio"] = t2_ratio - - _apply_ratio_overrides(ratio_dicts[0]) - _apply_ratio_overrides(ratio_dicts[1]) - apply_hidiffusion( model, apply_raunet=apply_raunet, apply_window_attn=apply_window_attn, + t1_ratio=t1_ratio, + t2_ratio=t2_ratio, has_controlnet=has_controlnet, is_controlnet_text_to_image=is_controlnet_text_to_image, generator=generator, + is_inpainting_task=is_inpainting_task, + use_aggressive_raunet=use_aggressive_raunet, + denoising_start=denoising_start, + denoising_end=denoising_end, ) yield finally: @@ -133,11 +114,6 @@ def _apply_ratio_overrides(ratio_dict: dict) -> None: except Exception as error: if not had_active_exception: teardown_error = error - if ratio_overrides is not None and ratio_dicts is not None: - ratio_dicts[0].clear() - ratio_dicts[0].update(ratio_overrides[0]) - ratio_dicts[1].clear() - ratio_dicts[1].update(ratio_overrides[1]) if original_num_upsamplers is not None: target.num_upsamplers = original_num_upsamplers if set_model_name_or_path: diff --git a/invokeai/frontend/web/openapi.json b/invokeai/frontend/web/openapi.json index ccd0f318694..199f5f89b6c 100644 --- a/invokeai/frontend/web/openapi.json +++ b/invokeai/frontend/web/openapi.json @@ -24636,28 +24636,42 @@ "type": "boolean" }, "hidiffusion_t1_ratio": { - "default": 0.4, + "anyOf": [ + { + "maximum": 1, + "minimum": 0, + "type": "number" + }, + { + "type": "null" + } + ], + "default": null, "description": "Override HiDiffusion early switch threshold (T1 ratio)", "field_kind": "input", "input": "any", - "maximum": 1, - "minimum": 0, - "orig_default": 0.4, + "orig_default": null, "orig_required": false, - "title": "HiDiffusion: T1 Ratio", - "type": "number" + "title": "HiDiffusion: T1 Ratio" }, "hidiffusion_t2_ratio": { - "default": 0.0, + "anyOf": [ + { + "maximum": 1, + "minimum": 0, + "type": "number" + }, + { + "type": "null" + } + ], + "default": null, "description": "Override HiDiffusion late switch threshold (T2 ratio)", "field_kind": "input", "input": "any", - "maximum": 1, - "minimum": 0, - "orig_default": 0.0, + "orig_default": null, "orig_required": false, - "title": "HiDiffusion: T2 Ratio", - "type": "number" + "title": "HiDiffusion: T2 Ratio" }, "latents": { "anyOf": [ @@ -24705,7 +24719,7 @@ "tags": ["latents", "denoise", "txt2img", "t2i", "t2l", "img2img", "i2i", "l2l"], "title": "Denoise - SD1.5, SDXL", "type": "object", - "version": "1.6.0", + "version": "1.7.0", "output": { "$ref": "#/components/schemas/LatentsOutput" } @@ -25049,28 +25063,42 @@ "type": "boolean" }, "hidiffusion_t1_ratio": { - "default": 0.4, + "anyOf": [ + { + "maximum": 1, + "minimum": 0, + "type": "number" + }, + { + "type": "null" + } + ], + "default": null, "description": "Override HiDiffusion early switch threshold (T1 ratio)", "field_kind": "input", "input": "any", - "maximum": 1, - "minimum": 0, - "orig_default": 0.4, + "orig_default": null, "orig_required": false, - "title": "HiDiffusion: T1 Ratio", - "type": "number" + "title": "HiDiffusion: T1 Ratio" }, "hidiffusion_t2_ratio": { - "default": 0.0, + "anyOf": [ + { + "maximum": 1, + "minimum": 0, + "type": "number" + }, + { + "type": "null" + } + ], + "default": null, "description": "Override HiDiffusion late switch threshold (T2 ratio)", "field_kind": "input", "input": "any", - "maximum": 1, - "minimum": 0, - "orig_default": 0.0, + "orig_default": null, "orig_required": false, - "title": "HiDiffusion: T2 Ratio", - "type": "number" + "title": "HiDiffusion: T2 Ratio" }, "latents": { "anyOf": [ @@ -25118,7 +25146,7 @@ "tags": ["latents", "denoise", "txt2img", "t2i", "t2l", "img2img", "i2i", "l2l"], "title": "Denoise - SD1.5, SDXL + Metadata", "type": "object", - "version": "1.2.0", + "version": "1.3.0", "output": { "$ref": "#/components/schemas/LatentsMetaOutput" } diff --git a/invokeai/frontend/web/public/locales/en.json b/invokeai/frontend/web/public/locales/en.json index c2a4853cc20..a63e37c9b89 100644 --- a/invokeai/frontend/web/public/locales/en.json +++ b/invokeai/frontend/web/public/locales/en.json @@ -1734,6 +1734,7 @@ "hiDiffusion": "HiDiffusion", "hiDiffusionRauNet": "HiDiffusion: RAU-Net", "hiDiffusionWindowAttn": "HiDiffusion: Window Attention", + "hiDiffusionRatiosAuto": "HiDiffusion: Automatic Ratios", "hiDiffusionT1Ratio": "HiDiffusion: T1 Ratio", "hiDiffusionT2Ratio": "HiDiffusion: T2 Ratio", "coherenceMode": "Mode", @@ -2183,6 +2184,13 @@ "Can boost local detail, but may affect global coherence." ] }, + "hidiffusionRatiosAuto": { + "heading": "HiDiffusion: Automatic Ratios", + "paragraphs": [ + "Uses the model- and resolution-specific T1 and T2 ratios recommended by HiDiffusion.", + "Disable this to set explicit ratio overrides below." + ] + }, "hidiffusionT1Ratio": { "heading": "HiDiffusion: T1 Ratio", "paragraphs": [ @@ -2194,7 +2202,7 @@ "heading": "HiDiffusion: T2 Ratio", "paragraphs": [ "Controls the late switch point for HiDiffusion (T2).", - "Higher values keep window attention active longer." + "Higher values keep the second RAU-Net resolution-aware stage active longer." ] }, "clipSkip": { diff --git a/invokeai/frontend/web/src/common/components/InformationalPopover/constants.ts b/invokeai/frontend/web/src/common/components/InformationalPopover/constants.ts index d91077f6607..edbe78ddab2 100644 --- a/invokeai/frontend/web/src/common/components/InformationalPopover/constants.ts +++ b/invokeai/frontend/web/src/common/components/InformationalPopover/constants.ts @@ -78,6 +78,7 @@ export type Feature = | 'hidiffusion' | 'hidiffusionRauNet' | 'hidiffusionWindowAttn' + | 'hidiffusionRatiosAuto' | 'hidiffusionT1Ratio' | 'hidiffusionT2Ratio' | 'colorCompensation' @@ -256,6 +257,7 @@ export const POPOVER_DATA: { [key in Feature]?: PopoverData } = { }, hidiffusionRauNet: {}, hidiffusionWindowAttn: {}, + hidiffusionRatiosAuto: {}, hidiffusionT1Ratio: {}, hidiffusionT2Ratio: {}, } as const; diff --git a/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.test.ts b/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.test.ts index e2a2713aa41..30c2690c72a 100644 --- a/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.test.ts +++ b/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.test.ts @@ -172,8 +172,9 @@ describe('paramsSlice selectors for external models', () => { * `animaT5EncoderModel`, since removed from the schema. * - v4: the narrowest v4 blob is not a release at all — it is the one written by the build that did * the bump, `1aeb05bbf0` (97 keys). Releases writing v4 start at v6.14.0-rc1. - * - v5: the current version, reached by the FLUX.2 [dev] merge `f10d2a4f5a`, which is also the - * build that wrote the narrowest v5 blob. Pinning the fixture at the bump commit is what keeps + * - v5: reached by the FLUX.2 [dev] merge `f10d2a4f5a`, which is also the build that wrote the + * narrowest v5 blob. The current v6 migration converts legacy HiDiffusion defaults to Auto. + * Pinning the fixture at the v5 bump commit is what keeps * the invariant below meaningful for the current tier: the version steps can never cover it (a * v5 blob matches no branch), so every key added since the bump has to carry a zod default, and * this entry is what proves it does. @@ -900,14 +901,14 @@ describe('paramsSliceConfig persisted state migration', () => { // v2 migrates all the way through the current chain (v2 -> v3 adds Qwen fields, // v3 -> v4 adds Krea-2 and PiD fields). - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.qwenImageVaeModel).toBeNull(); expect(result.qwenImageQwenVLEncoderModel).toBeNull(); expect(result.hiDiffusionEnabled).toBe(false); expect(result.hiDiffusionRauNetEnabled).toBe(true); expect(result.hiDiffusionWindowAttnEnabled).toBe(true); - expect(result.hiDiffusionT1Ratio).toBe(0.4); - expect(result.hiDiffusionT2Ratio).toBe(0.0); + expect(result.hiDiffusionT1Ratio).toBeNull(); + expect(result.hiDiffusionT2Ratio).toBeNull(); // Existing params should be preserved expect(result.positivePrompt).toBe('a fluffy cat'); expect(result.seed).toBe(42); @@ -916,6 +917,29 @@ describe('paramsSliceConfig persisted state migration', () => { expect(result.dimensions.height).toBe(768); }); + it('migrates old HiDiffusion defaults to automatic ratios and preserves custom overrides', () => { + expect(migrate).toBeDefined(); + const initial = getInitialParamsState(); + + const oldDefaults = migrate?.({ + ...initial, + _version: 5, + hiDiffusionT1Ratio: 0.4, + hiDiffusionT2Ratio: 0.0, + }) as ReturnType; + const customOverrides = migrate?.({ + ...initial, + _version: 5, + hiDiffusionT1Ratio: 0.65, + hiDiffusionT2Ratio: 0.25, + }) as ReturnType; + + expect(oldDefaults.hiDiffusionT1Ratio).toBeNull(); + expect(oldDefaults.hiDiffusionT2Ratio).toBeNull(); + expect(customOverrides.hiDiffusionT1Ratio).toBe(0.65); + expect(customOverrides.hiDiffusionT2Ratio).toBe(0.25); + }); + it('merges the separate Klein / dev VAE slots into flux2VaeModel when migrating from v3', () => { expect(migrate).toBeDefined(); @@ -938,7 +962,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(v3State) as ReturnType & Record; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect((result.flux2VaeModel as { key: string } | null)?.key).toBe('klein-vae'); // The new standalone dev Mistral encoder slot must be seeded, not left undefined. expect(result.flux2DevMistralEncoderModel).toBeNull(); @@ -972,7 +996,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(v3State) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.krea2VaeModel).toBeNull(); expect(result.krea2Qwen3VlEncoderModel).toBeNull(); expect(result.krea2SeedVarianceEnabled).toBe(false); @@ -1006,7 +1030,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(mainV4State) as ReturnType & Record; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect((result.flux2VaeModel as { key: string } | null)?.key).toBe('klein-vae'); expect(result.flux2DevMistralEncoderModel).toBeNull(); // main's own v4 values must survive untouched. @@ -1034,7 +1058,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(devV4State) as ReturnType & Record; - expect(result._version).toBe(5); + expect(result._version).toBe(6); // The branch's own v4 values must survive untouched. expect((result.flux2VaeModel as { key: string } | null)?.key).toBe('flux2-vae'); expect(result.pidMode).toBe('off'); @@ -1091,7 +1115,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(blob) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.positivePrompt).toBe('a fluffy cat'); expect(result.seed).toBe(42); expect(result.shouldRandomizeSeed).toBe(false); @@ -1139,7 +1163,7 @@ describe('paramsSliceConfig persisted state migration', () => { expect( backfilled, - version === getInitialParamsState()._version + Number(version) === getInitialParamsState()._version ? `Keys missing from a blob written at ${release}, the commit that bumped _version to ${version}. ` + `A blob already at the current version matches no branch in the migration chain, so no step can ` + `seed these — each needs a zod default, or upgrading throws in zParamsState.parse() and wipes ` + @@ -1257,7 +1281,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(blob) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.dimensions).toEqual(getInitialParamsState().dimensions); expect(result.positivePrompt).toBe('a fluffy cat'); expect(result.seed).toBe(7); @@ -1274,7 +1298,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(blob) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.positivePromptHistory).toEqual([]); expect(result.qwenImageVaeModel).toBeNull(); expect(result.wanVaeModel).toBeNull(); @@ -1284,18 +1308,18 @@ describe('paramsSliceConfig persisted state migration', () => { it('never repairs _version, so version detection cannot be bypassed', () => { // `_version` is the input to the version steps, so the net must leave it alone. If it repaired // it, any blob whose version is not the current literal — including one written by a *newer* - // build — would be silently stamped v5 having run no step, and its stale field values would be + // build — would be silently stamped with the current version having run no step, and its stale field values would be // accepted as current. Deliberately not routed through migrate(): the version steps normalise // `_version` before the net ever sees it, so only calling the net directly tests the guard. // The blob is otherwise complete (the current tier's key set), so `_version` is the only thing // the parse below can object to. - const blob = buildReleaseBlob('f10d2a4f5a', { _version: 6, positivePrompt: 'a fluffy cat' }); + const blob = buildReleaseBlob('f10d2a4f5a', { _version: 7, positivePrompt: 'a fluffy cat' }); const { backfilled, reset } = repairParamsState(blob); expect(backfilled).toEqual([]); expect(reset).toEqual([]); - expect(blob._version).toBe(6); + expect(blob._version).toBe(7); // Still fatal, which is the correct outcome for a downgrade: that slice really was written by a // schema this build does not know. expect(() => zParamsState.parse(blob)).toThrow(); @@ -1311,7 +1335,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(blob) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.positivePrompt).toBe('a fluffy cat'); expect(result.seed).toBe(7); expect(result.dimensions).toBeDefined(); @@ -1328,7 +1352,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(v3State) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.wanTransformerLowNoise).toBeNull(); expect(result.wanComponentSource).toBeNull(); expect(result.wanVaeModel).toBeNull(); @@ -1350,7 +1374,7 @@ describe('paramsSliceConfig persisted state migration', () => { const result = migrate?.(v2State) as ReturnType; - expect(result._version).toBe(5); + expect(result._version).toBe(6); expect(result.fluxScheduler).toBe('euler'); expect(result.zImageScheduler).toBe('euler'); expect(result.colorCompensation).toBe(false); @@ -1410,7 +1434,7 @@ describe('paramsSliceConfig persisted state migration', () => { expect('hiDiffusionEnabled' in blob).toBe(false); applyParamsVersionMigrations(blob); - expect(blob._version).toBe(5); + expect(blob._version).toBe(6); // The value assertions below cannot, on their own, prove the defaults exist: three mechanisms // produce the identical values, so any two can hide the third being reverted. Parsing directly @@ -1435,8 +1459,8 @@ describe('paramsSliceConfig persisted state migration', () => { expect(result.hiDiffusionEnabled).toBe(false); expect(result.hiDiffusionRauNetEnabled).toBe(true); expect(result.hiDiffusionWindowAttnEnabled).toBe(true); - expect(result.hiDiffusionT1Ratio).toBe(0.4); - expect(result.hiDiffusionT2Ratio).toBe(0.0); + expect(result.hiDiffusionT1Ratio).toBeNull(); + expect(result.hiDiffusionT2Ratio).toBeNull(); expect(result.positivePrompt).toBe('a fluffy cat'); expect(result.seed).toBe(42); }); diff --git a/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.ts b/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.ts index f7c8420dade..3db8e6b356d 100644 --- a/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.ts +++ b/invokeai/frontend/web/src/features/controlLayers/store/paramsSlice.ts @@ -182,10 +182,10 @@ const slice = createSlice({ setHiDiffusionWindowAttnEnabled: (state, action: PayloadAction) => { state.hiDiffusionWindowAttnEnabled = action.payload; }, - setHiDiffusionT1Ratio: (state, action: PayloadAction) => { + setHiDiffusionT1Ratio: (state, action: PayloadAction) => { state.hiDiffusionT1Ratio = action.payload; }, - setHiDiffusionT2Ratio: (state, action: PayloadAction) => { + setHiDiffusionT2Ratio: (state, action: PayloadAction) => { state.hiDiffusionT2Ratio = action.payload; }, setSeamlessXAxis: (state, action: PayloadAction) => { @@ -1162,6 +1162,15 @@ export const applyParamsVersionMigrations = (state: any): void => { state.gemma2EncoderModel = state.gemma2EncoderModel ?? null; state.pidSteps = state.pidSteps ?? 4; } + + if (state._version === 5) { + // v5 -> v6: numeric HiDiffusion defaults unintentionally overrode the library's + // resolution-aware presets. Treat the old default values as automatic thresholds while + // preserving values that users explicitly changed. + state._version = 6; + state.hiDiffusionT1Ratio = state.hiDiffusionT1Ratio === 0.4 ? null : state.hiDiffusionT1Ratio; + state.hiDiffusionT2Ratio = state.hiDiffusionT2Ratio === 0.0 ? null : state.hiDiffusionT2Ratio; + } }; export const paramsSliceConfig: SliceConfig = { diff --git a/invokeai/frontend/web/src/features/controlLayers/store/types.ts b/invokeai/frontend/web/src/features/controlLayers/store/types.ts index 2143e58f997..b6fff1580f7 100644 --- a/invokeai/frontend/web/src/features/controlLayers/store/types.ts +++ b/invokeai/frontend/web/src/features/controlLayers/store/types.ts @@ -817,7 +817,7 @@ const zPidMode = z.enum(['off', 'fit', 'native']); export type PidMode = z.infer; export const zParamsState = z.object({ - _version: z.literal(5), + _version: z.literal(6), maskBlur: z.number(), maskBlurMethod: zParameterMaskBlurMethod, canvasCoherenceMode: zParameterCanvasCoherenceMode, @@ -839,8 +839,8 @@ export const zParamsState = z.object({ hiDiffusionEnabled: z.boolean().default(false), hiDiffusionRauNetEnabled: z.boolean().default(true), hiDiffusionWindowAttnEnabled: z.boolean().default(true), - hiDiffusionT1Ratio: z.number().default(0.4), - hiDiffusionT2Ratio: z.number().default(0.0), + hiDiffusionT1Ratio: z.number().nullable().default(null), + hiDiffusionT2Ratio: z.number().nullable().default(null), iterations: z.number(), scheduler: zParameterScheduler, fluxScheduler: zParameterFluxScheduler, @@ -965,7 +965,7 @@ export const zParamsState = z.object({ }); export type ParamsState = z.infer; export const getInitialParamsState = (): ParamsState => ({ - _version: 5, + _version: 6, maskBlur: 16, maskBlurMethod: 'box', canvasCoherenceMode: 'Gaussian Blur', @@ -983,8 +983,8 @@ export const getInitialParamsState = (): ParamsState => ({ hiDiffusionEnabled: false, hiDiffusionRauNetEnabled: true, hiDiffusionWindowAttnEnabled: true, - hiDiffusionT1Ratio: 0.4, - hiDiffusionT2Ratio: 0.0, + hiDiffusionT1Ratio: null, + hiDiffusionT2Ratio: null, iterations: 1, scheduler: 'dpmpp_3m_k', fluxScheduler: 'euler', diff --git a/invokeai/frontend/web/src/features/metadata/parsing.test.ts b/invokeai/frontend/web/src/features/metadata/parsing.test.ts index 6072aa1c10f..12c8a96fa9e 100644 --- a/invokeai/frontend/web/src/features/metadata/parsing.test.ts +++ b/invokeai/frontend/web/src/features/metadata/parsing.test.ts @@ -1,5 +1,9 @@ import type { AppStore } from 'app/store/store'; -import { setHiDiffusionEnabled } from 'features/controlLayers/store/paramsSlice'; +import { + setHiDiffusionEnabled, + setHiDiffusionT1Ratio, + setHiDiffusionT2Ratio, +} from 'features/controlLayers/store/paramsSlice'; import { describe, expect, it, vi } from 'vitest'; import { ImageMetadataHandlers, MetadataUtils, parseMetadataHandler } from './parsing'; @@ -163,12 +167,31 @@ describe('Qwen metadata parsing', () => { }); describe('HiDiffusion metadata parsing', () => { + it('recalls null ratios as automatic thresholds', async () => { + const store = createStore(); + const metadata = { hidiffusion_t1_ratio: null, hidiffusion_t2_ratio: null }; + + const t1 = await parseMetadataHandler(metadata, ImageMetadataHandlers.HiDiffusionT1Ratio, store); + const t2 = await parseMetadataHandler(metadata, ImageMetadataHandlers.HiDiffusionT2Ratio, store); + ImageMetadataHandlers.HiDiffusionT1Ratio.recall(t1, store); + ImageMetadataHandlers.HiDiffusionT2Ratio.recall(t2, store); + + expect(store.dispatch).toHaveBeenCalledWith(setHiDiffusionT1Ratio(null)); + expect(store.dispatch).toHaveBeenCalledWith(setHiDiffusionT2Ratio(null)); + }); + it('disables HiDiffusion when recalling all metadata from an older image', async () => { let hiDiffusionEnabled = true; + let hiDiffusionT1Ratio: number | null = 0.8; + let hiDiffusionT2Ratio: number | null = 0.6; const store = { dispatch: vi.fn((action) => { if (action.type === setHiDiffusionEnabled.type) { hiDiffusionEnabled = action.payload; + } else if (action.type === setHiDiffusionT1Ratio.type) { + hiDiffusionT1Ratio = action.payload; + } else if (action.type === setHiDiffusionT2Ratio.type) { + hiDiffusionT2Ratio = action.payload; } return action; }), @@ -193,5 +216,7 @@ describe('HiDiffusion metadata parsing', () => { expect(store.dispatch).toHaveBeenCalledWith(setHiDiffusionEnabled(false)); expect(hiDiffusionEnabled).toBe(false); + expect(hiDiffusionT1Ratio).toBeNull(); + expect(hiDiffusionT2Ratio).toBeNull(); }); }); diff --git a/invokeai/frontend/web/src/features/metadata/parsing.tsx b/invokeai/frontend/web/src/features/metadata/parsing.tsx index a653afc10b8..0c57a4058d9 100644 --- a/invokeai/frontend/web/src/features/metadata/parsing.tsx +++ b/invokeai/frontend/web/src/features/metadata/parsing.tsx @@ -787,12 +787,17 @@ const HiDiffusionWindowAttn: SingleMetadataHandler = { //#endregion HiDiffusionWindowAttn //#region HiDiffusionT1Ratio -const HiDiffusionT1Ratio: SingleMetadataHandler = { +const HiDiffusionRatioValue = ({ value }: SingleMetadataValueProps) => { + const { t } = useTranslation(); + return ; +}; + +const HiDiffusionT1Ratio: SingleMetadataHandler = { [SingleMetadataKey]: true, type: 'HiDiffusionT1Ratio', parse: (metadata, _store) => { const raw = getProperty(metadata, 'hidiffusion_t1_ratio'); - const parsed = z.number().parse(raw); + const parsed = raw === undefined ? null : z.number().nullable().parse(raw); return Promise.resolve(parsed); }, recall: (value, store) => { @@ -800,17 +805,17 @@ const HiDiffusionT1Ratio: SingleMetadataHandler = { }, i18nKey: 'metadata.hiDiffusionT1Ratio', LabelComponent: MetadataLabel, - ValueComponent: ({ value }: SingleMetadataValueProps) => , + ValueComponent: HiDiffusionRatioValue, }; //#endregion HiDiffusionT1Ratio //#region HiDiffusionT2Ratio -const HiDiffusionT2Ratio: SingleMetadataHandler = { +const HiDiffusionT2Ratio: SingleMetadataHandler = { [SingleMetadataKey]: true, type: 'HiDiffusionT2Ratio', parse: (metadata, _store) => { const raw = getProperty(metadata, 'hidiffusion_t2_ratio'); - const parsed = z.number().parse(raw); + const parsed = raw === undefined ? null : z.number().nullable().parse(raw); return Promise.resolve(parsed); }, recall: (value, store) => { @@ -818,7 +823,7 @@ const HiDiffusionT2Ratio: SingleMetadataHandler = { }, i18nKey: 'metadata.hiDiffusionT2Ratio', LabelComponent: MetadataLabel, - ValueComponent: ({ value }: SingleMetadataValueProps) => , + ValueComponent: HiDiffusionRatioValue, }; //#endregion HiDiffusionT2Ratio diff --git a/invokeai/frontend/web/src/features/nodes/util/graph/generation/addSDXLRefiner.ts b/invokeai/frontend/web/src/features/nodes/util/graph/generation/addSDXLRefiner.ts index 11a5333c7d2..2f755377df9 100644 --- a/invokeai/frontend/web/src/features/nodes/util/graph/generation/addSDXLRefiner.ts +++ b/invokeai/frontend/web/src/features/nodes/util/graph/generation/addSDXLRefiner.ts @@ -65,8 +65,8 @@ export const addSDXLRefiner = async ( hidiffusion: hiDiffusionEnabled, hidiffusion_raunet: hiDiffusionRauNetEnabled, hidiffusion_window_attn: hiDiffusionWindowAttnEnabled, - hidiffusion_t1_ratio: hiDiffusionEnabled ? hiDiffusionT1Ratio : undefined, - hidiffusion_t2_ratio: hiDiffusionEnabled ? hiDiffusionT2Ratio : undefined, + hidiffusion_t1_ratio: hiDiffusionEnabled ? (hiDiffusionT1Ratio ?? undefined) : undefined, + hidiffusion_t2_ratio: hiDiffusionEnabled ? (hiDiffusionT2Ratio ?? undefined) : undefined, denoising_start: refinerStart, denoising_end: 1, }); diff --git a/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildHiDiffusionGraph.test.ts b/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildHiDiffusionGraph.test.ts index 8853af97baa..6ed9b077a82 100644 --- a/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildHiDiffusionGraph.test.ts +++ b/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildHiDiffusionGraph.test.ts @@ -32,8 +32,8 @@ const defaultParams = { cfgRescaleMultiplier: 0, hiDiffusionEnabled: false, hiDiffusionRauNetEnabled: false, - hiDiffusionT1Ratio: 0.25, - hiDiffusionT2Ratio: 0.1, + hiDiffusionT1Ratio: 0.25 as number | null, + hiDiffusionT2Ratio: 0.1 as number | null, hiDiffusionWindowAttnEnabled: false, scheduler: 'euler', steps: 20, @@ -193,4 +193,21 @@ describe('HiDiffusion graph metadata', () => { expect(metadata.hidiffusion_t1_ratio).toBe(0.25); expect(metadata.hidiffusion_t2_ratio).toBe(0.1); }); + + it('omits automatic ratio overrides from the SDXL denoise and metadata nodes', async () => { + currentModel = sdxlModel; + params.hiDiffusionEnabled = true; + params.hiDiffusionRauNetEnabled = true; + params.hiDiffusionT1Ratio = null; + params.hiDiffusionT2Ratio = null; + + const { g } = await buildSDXLGraph(buildGraphArg()); + const denoise = g.getNodes().find((node) => node.type === 'denoise_latents'); + const metadata = getMetadata(g); + + expect(denoise?.hidiffusion_t1_ratio).toBeUndefined(); + expect(denoise?.hidiffusion_t2_ratio).toBeUndefined(); + expect(metadata.hidiffusion_t1_ratio).toBeNull(); + expect(metadata.hidiffusion_t2_ratio).toBeNull(); + }); }); diff --git a/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildSD1Graph.ts b/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildSD1Graph.ts index 2a8d15b1e75..8ac62a08885 100644 --- a/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildSD1Graph.ts +++ b/invokeai/frontend/web/src/features/nodes/util/graph/generation/buildSD1Graph.ts @@ -109,8 +109,8 @@ export const buildSD1Graph = async (arg: GraphBuilderArg): Promise { + const hiDiffusionEnabled = useAppSelector(selectHiDiffusionEnabled); + const hiDiffusionRauNetEnabled = useAppSelector(selectHiDiffusionRauNetEnabled); + const hiDiffusionT1Ratio = useAppSelector(selectHiDiffusionT1Ratio); + const hiDiffusionT2Ratio = useAppSelector(selectHiDiffusionT2Ratio); + const dispatch = useAppDispatch(); + const { t } = useTranslation(); + const isAutomatic = hiDiffusionT1Ratio === null && hiDiffusionT2Ratio === null; + + const onChange = useCallback( + (event: ChangeEvent) => { + if (event.target.checked) { + dispatch(setHiDiffusionT1Ratio(null)); + dispatch(setHiDiffusionT2Ratio(null)); + } else { + dispatch(setHiDiffusionT1Ratio(RATIO_CONSTRAINTS.t1.initial)); + dispatch(setHiDiffusionT2Ratio(RATIO_CONSTRAINTS.t2.initial)); + } + }, + [dispatch] + ); + + return ( + + + + {t('parameters.hiDiffusionRatiosAuto')} + + + + + ); +}); + +ParamHiDiffusionAutoRatiosToggle.displayName = 'ParamHiDiffusionAutoRatiosToggle'; + export const ParamHiDiffusionT1Ratio = memo(() => { const hiDiffusionEnabled = useAppSelector(selectHiDiffusionEnabled); + const hiDiffusionRauNetEnabled = useAppSelector(selectHiDiffusionRauNetEnabled); const hiDiffusionT1Ratio = useAppSelector(selectHiDiffusionT1Ratio); + const hiDiffusionT2Ratio = useAppSelector(selectHiDiffusionT2Ratio); const dispatch = useAppDispatch(); const { t } = useTranslation(); const onChange = useCallback((value: number) => dispatch(setHiDiffusionT1Ratio(value)), [dispatch]); return ( - + {t('parameters.hiDiffusionT1Ratio')} { marks /> { const hiDiffusionEnabled = useAppSelector(selectHiDiffusionEnabled); + const hiDiffusionRauNetEnabled = useAppSelector(selectHiDiffusionRauNetEnabled); + const hiDiffusionT1Ratio = useAppSelector(selectHiDiffusionT1Ratio); const hiDiffusionT2Ratio = useAppSelector(selectHiDiffusionT2Ratio); const dispatch = useAppDispatch(); const { t } = useTranslation(); @@ -167,14 +216,19 @@ export const ParamHiDiffusionT2Ratio = memo(() => { const onChange = useCallback((value: number) => dispatch(setHiDiffusionT2Ratio(value)), [dispatch]); return ( - + {t('parameters.hiDiffusionT2Ratio')} { marks /> { + diff --git a/invokeai/frontend/web/src/services/api/schema.ts b/invokeai/frontend/web/src/services/api/schema.ts index c19832d09a1..d2b880fb4ec 100644 --- a/invokeai/frontend/web/src/services/api/schema.ts +++ b/invokeai/frontend/web/src/services/api/schema.ts @@ -9338,15 +9338,15 @@ export type components = { /** * HiDiffusion: T1 Ratio * @description Override HiDiffusion early switch threshold (T1 ratio) - * @default 0.4 + * @default null */ - hidiffusion_t1_ratio?: number; + hidiffusion_t1_ratio?: number | null; /** * HiDiffusion: T2 Ratio * @description Override HiDiffusion late switch threshold (T2 ratio) - * @default 0 + * @default null */ - hidiffusion_t2_ratio?: number; + hidiffusion_t2_ratio?: number | null; /** * @description Latents tensor * @default null @@ -9486,15 +9486,15 @@ export type components = { /** * HiDiffusion: T1 Ratio * @description Override HiDiffusion early switch threshold (T1 ratio) - * @default 0.4 + * @default null */ - hidiffusion_t1_ratio?: number; + hidiffusion_t1_ratio?: number | null; /** * HiDiffusion: T2 Ratio * @description Override HiDiffusion late switch threshold (T2 ratio) - * @default 0 + * @default null */ - hidiffusion_t2_ratio?: number; + hidiffusion_t2_ratio?: number | null; /** * @description Latents tensor * @default null diff --git a/tests/backend/stable_diffusion/test_hidiffusion_utils.py b/tests/backend/stable_diffusion/test_hidiffusion_utils.py index 5f8619a4882..b7e9bad9130 100644 --- a/tests/backend/stable_diffusion/test_hidiffusion_utils.py +++ b/tests/backend/stable_diffusion/test_hidiffusion_utils.py @@ -6,7 +6,10 @@ import torch from invokeai.backend.hidiffusion.hidiffusion import ( + _get_switching_threshold_ratio, _resize_controlnet_residual, + _should_use_aggressive_raunet, + make_diffusers_downsampler_block, switching_threshold_ratio_dict, text_to_img_controlnet_switching_threshold_ratio_dict, ) @@ -256,16 +259,14 @@ def test_hidiffusion_patch_restores_state_when_apply_hidiffusion_raises(): ) hook = MagicMock() - def fake_apply_hidiffusion(patched_model, **_kwargs): + def fake_apply_hidiffusion(patched_model, **kwargs): assert patched_model._name_or_path == "patched-model-name" assert patched_model.config._name_or_path == "patched-model-name" - first_switching_entry = next(iter(switching_threshold_ratio_dict.values())) - first_controlnet_entry = next(iter(text_to_img_controlnet_switching_threshold_ratio_dict.values())) - assert first_switching_entry["T1_ratio"] == 0.25 - assert first_switching_entry["T2_ratio"] == 0.1 - assert first_controlnet_entry["T1_ratio"] == 0.25 - assert first_controlnet_entry["T2_ratio"] == 0.1 + assert kwargs["t1_ratio"] == 0.25 + assert kwargs["t2_ratio"] == 0.1 + assert switching_threshold_ratio_dict == original_switching + assert text_to_img_controlnet_switching_threshold_ratio_dict == original_controlnet patched_model.unet.num_upsamplers = 99 patched_model.unet.layer.info = {"hooks": [hook]} @@ -369,3 +370,144 @@ def __getattr__(self, name): assert config._internal_dict["_name_or_path"] == "patched-model-name" assert "_name_or_path" not in config._internal_dict + + +def test_hidiffusion_ratio_overrides_are_isolated_between_overlapping_patches(): + original_switching = copy.deepcopy(switching_threshold_ratio_dict) + original_controlnet = copy.deepcopy(text_to_img_controlnet_switching_threshold_ratio_dict) + first_model = SimpleNamespace(unet=DummyUNet()) + second_model = SimpleNamespace(unet=DummyUNet()) + applied_overrides: list[tuple[object, float | None, float | None]] = [] + + def fake_apply_hidiffusion(model, **kwargs): + applied_overrides.append((model, kwargs["t1_ratio"], kwargs["t2_ratio"])) + + with ( + patch("invokeai.backend.hidiffusion.hidiffusion.apply_hidiffusion", side_effect=fake_apply_hidiffusion), + patch("invokeai.backend.hidiffusion.hidiffusion.remove_hidiffusion"), + ): + first_patch = hidiffusion_patch(first_model, name_or_path="first", t1_ratio=0.2, t2_ratio=0.1) + second_patch = hidiffusion_patch(second_model, name_or_path="second", t1_ratio=0.8, t2_ratio=0.9) + first_patch.__enter__() + second_patch.__enter__() + first_patch.__exit__(None, None, None) + second_patch.__exit__(None, None, None) + + assert applied_overrides == [(first_model, 0.2, 0.1), (second_model, 0.8, 0.9)] + assert switching_threshold_ratio_dict == original_switching + assert text_to_img_controlnet_switching_threshold_ratio_dict == original_controlnet + + +def test_hidiffusion_patch_forwards_generation_context(): + model = SimpleNamespace(unet=DummyUNet()) + + with ( + patch("invokeai.backend.hidiffusion.hidiffusion.apply_hidiffusion") as mock_apply_hidiffusion, + patch("invokeai.backend.hidiffusion.hidiffusion.remove_hidiffusion"), + ): + with hidiffusion_patch( + model, + name_or_path="stabilityai/stable-diffusion-xl-base-1.0", + is_inpainting_task=True, + use_aggressive_raunet=False, + denoising_start=0.6, + denoising_end=1.0, + ): + pass + + kwargs = mock_apply_hidiffusion.call_args.kwargs + assert kwargs["is_inpainting_task"] is True + assert kwargs["use_aggressive_raunet"] is False + assert kwargs["denoising_start"] == 0.6 + assert kwargs["denoising_end"] == 1.0 + + +@pytest.mark.parametrize( + ("denoising_start", "denoising_end", "expected_ratio"), + [ + (0.0, 1.0, 0.4), + (0.6, 1.0, 0.0), + (0.0, 0.2, 1.0), + (0.2, 0.6, 0.5), + ], +) +def test_hidiffusion_ratios_are_mapped_to_the_executed_denoising_range( + denoising_start: float, denoising_end: float, expected_ratio: float +): + module = SimpleNamespace( + switching_threshold_ratio="T1_ratio", + info={ + "switching_threshold_overrides": {"T1_ratio": None, "T2_ratio": None}, + "denoising_start": denoising_start, + "denoising_end": denoising_end, + }, + ) + + ratio = _get_switching_threshold_ratio(module, switching_threshold_ratio_dict, "sdxl_2048") + + assert ratio == pytest.approx(expected_ratio) + + +@pytest.mark.parametrize( + ("is_inpainting_task", "override", "expected"), + [ + (False, False, False), + (True, True, True), + (True, None, False), + (False, None, True), + ], +) +def test_explicit_aggressive_raunet_setting_takes_precedence( + is_inpainting_task: bool, override: bool | None, expected: bool +): + module = SimpleNamespace( + info={ + "is_inpainting_task": is_inpainting_task, + "is_playground": False, + "use_aggressive_raunet": override, + } + ) + + assert _should_use_aggressive_raunet(module) is expected + + +def test_sdxl_t2_override_controls_downsampler_at_2048_resolution(): + patched_conv = make_diffusers_downsampler_block(torch.nn.Conv2d) + hidden_states = torch.arange(64, dtype=torch.float32).reshape(1, 1, 8, 8) + + def run(t2_ratio: float) -> torch.Tensor: + module = patched_conv(1, 1, kernel_size=3, stride=2, padding=1, bias=False) + torch.nn.init.constant_(module.weight, 1.0) + module.info = { + "size": (256, 256), + "pipeline": SimpleNamespace(_num_timesteps=30), + "text_to_img_controlnet": False, + "is_inpainting_task": False, + "is_playground": False, + "switching_threshold_overrides": {"T1_ratio": None, "T2_ratio": t2_ratio}, + } + module.model = "sdxl" + module.switching_threshold_ratio = "T2_ratio" + return module(hidden_states) + + assert not torch.equal(run(0.0), run(1.0)) + + +def test_sdxl_automatic_ratios_preserve_extreme_resolution_preset(): + patched_conv = make_diffusers_downsampler_block(torch.nn.Conv2d) + module = patched_conv(1, 1, kernel_size=3, stride=2, padding=1, bias=False) + module.info = { + "size": (512, 512), + "pipeline": SimpleNamespace(_num_timesteps=30), + "text_to_img_controlnet": False, + "is_inpainting_task": False, + "is_playground": False, + "switching_threshold_overrides": {"T1_ratio": None, "T2_ratio": None}, + } + module.model = "sdxl" + module.switching_threshold_ratio = "T2_ratio" + + module(torch.ones(1, 1, 8, 8)) + + assert module.T1_ratio == 0.3 + assert module.T1 == 9 diff --git a/tests/backend/stable_diffusion/test_regional_ip_data.py b/tests/backend/stable_diffusion/test_regional_ip_data.py new file mode 100644 index 00000000000..734fac305e3 --- /dev/null +++ b/tests/backend/stable_diffusion/test_regional_ip_data.py @@ -0,0 +1,18 @@ +import torch + +from invokeai.backend.stable_diffusion.diffusion.regional_ip_data import RegionalIPData + + +def test_regional_ip_data_supports_hidiffusion_raunet_downscale() -> None: + regional_ip_data = RegionalIPData( + image_prompt_embeds=[torch.zeros((1, 1, 4, 8))], + scales=[1.0], + masks=[torch.ones((1, 1, 192, 192))], + dtype=torch.float32, + device=torch.device("cpu"), + ) + + masks = regional_ip_data.get_masks(query_seq_len=12 * 12) + + assert masks.shape == (1, 1, 12 * 12, 1) + assert torch.count_nonzero(masks) == masks.numel() diff --git a/tests/backend/stable_diffusion/test_regional_prompt_data.py b/tests/backend/stable_diffusion/test_regional_prompt_data.py new file mode 100644 index 00000000000..e1990838c2c --- /dev/null +++ b/tests/backend/stable_diffusion/test_regional_prompt_data.py @@ -0,0 +1,21 @@ +import torch + +from invokeai.backend.stable_diffusion.diffusion.conditioning_data import Range, TextConditioningRegions +from invokeai.backend.stable_diffusion.diffusion.regional_prompt_data import RegionalPromptData + + +def test_regional_prompt_data_supports_hidiffusion_raunet_downscale() -> None: + regions = TextConditioningRegions( + masks=torch.ones((1, 1, 192, 192), dtype=torch.bool), + ranges=[Range(start=0, end=4)], + ) + regional_prompt_data = RegionalPromptData( + regions=[regions], + device=torch.device("cpu"), + dtype=torch.float32, + ) + + attention_mask = regional_prompt_data.get_cross_attn_mask(query_seq_len=12 * 12, key_seq_len=4) + + assert attention_mask.shape == (1, 12 * 12, 4) + assert torch.count_nonzero(attention_mask) == 0